From b253f672c07e73ec9fac21b0df6271e77d62533c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 06:12:00 +0000 Subject: [PATCH 01/40] Price the recomputed layer's attention or GDN activations in the checkpoint floor Full one-layer recompute replays a layer with gradients, so its mixer's saved activations stay live beside that layer's MoE stage. The checkpoint floor priced boundaries and the MoE stage only, which left context-parallel runs short: Qwen3.6-35B-A3B at CP2 peaked 9-11 GB above the floor on the most loaded rank. Price the larger of the model's attention and GDN mixers per recomputed row, with context-parallel stage buffers and GDN exchange copies, from allocator traces at CP1 and CP2. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 107 +++++++++++++----- .../test_trainer_rank_checkpoint_memory.py | 64 ++++++++++- .../test_trainer_rank_converted_memory.py | 8 +- 3 files changed, 147 insertions(+), 32 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 572ad7b67..fe45a9ba9 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4024,15 +4024,18 @@ def _checkpoint_memory_floor( group_rows: tuple[tuple[int, bool], ...], slot_refs: tuple["LoRASlotRef | None", ...] | None = None, ) -> tuple[int, int]: - """Conservative saved-boundary charge and one disjoint MoE workspace. + """Conservative saved-boundary charge and one recomputed layer's workspace. Count actual local full/uniform/1 boundaries, including aliases, rather than claiming measured distinct storage. Only this call's new groups enter the term; already-live graphs remain in the availability baseline. - No-grad groups also keep decoder input, current layer input, its MLP - residual and norm output across the MoE stage. Count these four row - tensors separately from returned outputs, allowing storage aliases. - This is not a bound for custom preprocessing, attention, or all backward. + 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) if not group_rows or len(self.runtime.model) != 1: @@ -4091,9 +4094,10 @@ def _checkpoint_memory_floor( if gradient_rows: self._checkpoint_moe_bytes_per_token() refs = (None,) * len(group_rows) if slot_refs is None else slot_refs + mixer = self._recomputed_mixer_bytes_per_token() if gradient_rows else 0 workspace = max( self._moe_workspace_bytes(rows, checkpoint_grad=grad, slot_ref=ref) - + (0 if grad else 4 * rows * self._hidden_size * 2) + + (mixer * rows if grad else 4 * rows * self._hidden_size * 2) for (rows, grad), ref in zip(group_rows, refs, strict=True) ) return retained, workspace @@ -7834,28 +7838,7 @@ def _estimate_required_memory_bytes_from_values( # Gathered LoRA inputs alias norm output without sequence sharding. gathered = hidden if sp > 1 else 0 common = 2 * hidden / sp + gathered - attention_width = ( - geometry.num_attention_heads * geometry.kv_channels or hidden - ) - kv_width = geometry.num_query_groups * geometry.kv_channels or hidden - gated = self._attention_output_gate - attention = ( - common + ((7 if gated else 5) * attention_width + 3 * kv_width) / tp - ) - if 0 < geometry.num_query_groups < tp: - # SelfAttentionLinearQKVLoRA constructs global QKV before - # slicing it when KV groups cannot be partitioned across TP. - attention += ((2 if gated else 1) * attention_width + 2 * kv_width) * ( - 1 - 1 / tp - ) - gdn = ( - common - + ( - 4 * geometry.gdn_key_heads * geometry.gdn_key_head_dim - + 8 * geometry.gdn_value_heads * geometry.gdn_value_head_dim - ) - / tp - ) + attention, gdn = self._mixer_activation_widths() ffn_width = geometry.ffn_hidden_size or 4 * hidden mlp = common + self._mlp_activation_factor * ffn_width / tp if geometry.moe_experts: @@ -7998,6 +7981,74 @@ def active(chunk: torch.nn.Module) -> bool: return all(active(chunk) for chunk in self.runtime.model) + def _mixer_activation_widths(self) -> tuple[float, float]: + """Saved activations per token of one attention and one GDN layer. + + Elements, not bytes: what a layer's attention or GDN mixer keeps for + its backward, including the input norm output (and gathered LoRA + inputs under sequence parallelism). + """ + geometry = self._geometry + hidden = self._hidden_size + tp = max(1, self._topology_key()[1]) + sp = tp if self._sequence_parallel else 1 + # Gathered LoRA inputs alias norm output without sequence sharding. + gathered = hidden if sp > 1 else 0 + common = 2 * hidden / sp + gathered + attention_width = geometry.num_attention_heads * geometry.kv_channels or hidden + kv_width = geometry.num_query_groups * geometry.kv_channels or hidden + gated = self._attention_output_gate + attention = common + ((7 if gated else 5) * attention_width + 3 * kv_width) / tp + if 0 < geometry.num_query_groups < tp: + # SelfAttentionLinearQKVLoRA constructs global QKV before + # slicing it when KV groups cannot be partitioned across TP. + attention += ((2 if gated else 1) * attention_width + 2 * kv_width) * ( + 1 - 1 / tp + ) + gdn = ( + common + + ( + 4 * geometry.gdn_key_heads * geometry.gdn_key_head_dim + + 8 * geometry.gdn_value_heads * geometry.gdn_value_head_dim + ) + / tp + ) + return attention, gdn + + def _recomputed_mixer_bytes_per_token(self) -> int: + """Saved mixer activations of the layer being recomputed, per row. + + Full recompute replays one layer with gradients, and its attention or + GDN mixer keeps what its backward needs across that layer's MoE stage. + Price the larger mixer the model has, from Qwen3.6-35B-A3B allocator + traces on H200 (bytes per local token, at the layer's recompute peak): + + - attention: 66 KB at CP1, the retained width below. CP ranks also + keep stage-padded Q/K/V, the stage output and a core-attention copy: + 94 KB at CP2, priced at 95 KB. + - GDN: 75 KB at CP1, fewer tensors than the retained width, priced at + 74 KB. CP ranks add rank-exchange copies: 84 KB at CP2, priced at + 86 KB. + """ + geometry = self._geometry + hidden = self._hidden_size + cp = self._topology_key()[2] > 1 + widths = [] + if self._gdn_layers < self._num_layers: + attention, _gdn = self._mixer_activation_widths() + if cp: + q = geometry.num_attention_heads * geometry.kv_channels or hidden + kv = geometry.num_query_groups * geometry.kv_channels or hidden + attention += 3 * q + 2 * kv + widths.append(attention) + if self._gdn_layers: + key = geometry.gdn_key_heads * geometry.gdn_key_head_dim + value = geometry.gdn_value_heads * geometry.gdn_value_head_dim + widths.append( + 2 * hidden + 4 * key + 6 * value + ((key + value) if cp else 0) + ) + return int(max(widths, default=0) * self._param_dtype_size) + def _gdn_segment_layer_bytes(self) -> float: """Initial and final fp32 recurrent states plus convolution history.""" geometry = self._geometry diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index c6e33ee9d..8542f5ec0 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -193,11 +193,18 @@ def test_topology_revalidated(axis): @pytest.mark.parametrize("rows", [(10, True), (11, False)]) @pytest.mark.parametrize("cp", [2, 4]) def test_cp_floor_prices_rank_rows(cp, rows): - # Callers pass rows on the most loaded CP rank; the per-row floor matches CP1. + # Callers pass rows on the most loaded CP rank; the per-row floor matches + # CP1, except that CP attention also keeps its stage buffers while a + # gradient group recomputes the layer. r = rank() single = r._checkpoint_memory_floor((rows,)) r._topology_key = lambda: (1, 1, cp, 1) - assert r._checkpoint_memory_floor((rows,)) == single != (0, 0) + retained, workspace = r._checkpoint_memory_floor((rows,)) + assert retained == single[0] and single != (0, 0) + count, grad = rows + # No attention geometry in this stub: Q and KV widths fall back to hidden. + stage = count * 2 * (3 * 2048 + 2 * 2048) if grad else 0 + assert workspace == single[1] + stage @pytest.mark.parametrize("share", [lambda n: -(-n // 2), lambda n: n * 3 // 4]) @@ -497,3 +504,56 @@ def test_no_grad_enclosure_config_guard(field, value): r = rank() setattr(r.runtime.model[0].decoder.config, field, value) assert r._checkpoint_memory_floor(((11, False),)) == (0, 0) + + +def qwen36_attention(r): + # Qwen3.6-35B-A3B attention: 16 heads and 2 query groups of 256, gated. + r._geometry = replace( + r._geometry, num_attention_heads=16, num_query_groups=2, kv_channels=256 + ) + r._attention_output_gate = True + return r + + +@pytest.mark.parametrize( + "cp,expected", + # Traces measured 66 KB per token at CP1 and 94 KB at CP2. + [(1, 2 * (2 * 2048 + 7 * 4096 + 3 * 512)), (2, 95232), (4, 95232)], +) +def test_recomputed_attention_is_priced_beside_the_moe_stage(cp, expected): + r = qwen36_attention(rank()) + r._topology_key = lambda: (1, 1, cp, 1) + assert r._recomputed_mixer_bytes_per_token() == expected + moe = r._moe_workspace_bytes(10, checkpoint_grad=True) + assert r._checkpoint_memory_floor(((10, True),))[1] == moe + 10 * expected + + +@pytest.mark.parametrize( + "cp,gdn_layers,expected", + [ + # Hybrid: GDN (74 KB) is larger at CP1, CP attention (95 KB) at CP2. + (1, 30, 73728), + (2, 30, 95232), + # GDN only: its CP rank-exchange copies add a key and a value row. + (1, 40, 73728), + (2, 40, 86016), + ], +) +def test_larger_recomputed_mixer_is_priced(cp, gdn_layers, expected): + r = qwen36_attention(rank()) + r._geometry = replace( + r._geometry, + gdn_key_heads=16, + gdn_key_head_dim=128, + gdn_value_heads=32, + gdn_value_head_dim=128, + ) + r._gdn_layers = gdn_layers + r._topology_key = lambda: (1, 1, cp, 1) + assert r._recomputed_mixer_bytes_per_token() == expected + + +def test_no_grad_groups_do_not_recompute_a_mixer(): + r = qwen36_attention(rank()) + moe = r._moe_workspace_bytes(10) + assert r._checkpoint_memory_floor(((10, False),)) == (0, moe + 4 * 10 * 2048 * 2) diff --git a/tests/unit/test_trainer_rank_converted_memory.py b/tests/unit/test_trainer_rank_converted_memory.py index a9c809163..6e8adf9bb 100644 --- a/tests/unit/test_trainer_rank_converted_memory.py +++ b/tests/unit/test_trainer_rank_converted_memory.py @@ -96,7 +96,10 @@ def test_actual_plan_cost_and_admission(layer, rank_value, grad, output): retained, workspace = rank._checkpoint_memory_floor(rank._plan_group_rows(plan)) pending = _gdn_memory.plan_floor(rank, plan) if grad: - assert workspace == expected(8, rank_value, True) + # The recomputed layer's mixer stays live beside its MoE stage. + mixer = rank._recomputed_mixer_bytes_per_token() + assert mixer > 0 + assert workspace == expected(8, rank_value, True) + 8 * mixer assert pending[0] == retained == 8 * 40 * 2048 * 2 assert pending[1] >= workspace else: @@ -125,7 +128,8 @@ def test_reference_and_gradient_keep_distinct_stage_modes(layer, order, rank_val retained, workspace = rank._checkpoint_memory_floor(groups) assert retained == 3 * 40 * 2048 * 2 assert workspace == max( - expected(3, rank_value, True), expected(9, rank_value, False) + 4 * 9 * 2048 * 2 + expected(3, rank_value, True) + 3 * rank._recomputed_mixer_bytes_per_token(), + expected(9, rank_value, False) + 4 * 9 * 2048 * 2, ) assert ( rank._memory_check(plan).estimated_required_bytes From 7cf9aa0c7b8b89a863fba7fb86ac0eb7d0ceb263 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 06:26:47 +0000 Subject: [PATCH 02/40] Count the recomputed mixer in generic admission without a MoE component Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_trainer_rank_pending_memory.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/tests/unit/test_trainer_rank_pending_memory.py b/tests/unit/test_trainer_rank_pending_memory.py index 490352f76..7fe3cf5eb 100644 --- a/tests/unit/test_trainer_rank_pending_memory.py +++ b/tests/unit/test_trainer_rank_pending_memory.py @@ -340,9 +340,11 @@ def test_constructor_declined_moe_keeps_generic_admission(layer, unsupported): plan = rank._plan_flat_forward(requests) assert g.plan_floor(rank, plan) == (0, 0) required = rank._plan_cost(plan).required - # Generic checkpoint-input accounting still applies without a MoE component. + # Generic checkpoint accounting, including the recomputed layer's mixer, + # still applies without a MoE component. gradient = 50640 * 40 * 2048 * 2 - assert required == int((plan.output_bytes + 2 * gradient) * 1.1) + mixer = 50640 * rank._recomputed_mixer_bytes_per_token() + assert required == int((plan.output_bytes + 2 * gradient + mixer) * 1.1) rank._available_memory_bytes = lambda: required - 1 assert not rank._memory_check(plan).fits rank._available_memory_bytes = lambda: required From a256cd3c7666c9c8a349d177b907aea2aebda9db Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 06:57:32 +0000 Subject: [PATCH 03/40] Ground the GDN recompute width in its measured sites and exchange widths Price GDN from its saved tensors (norm output, q/k with fp32 l2norm copies, v, z, segment-layout tensors, gated norm and the chunk decay matrix) instead of a ratio fit, and its context-parallel exchanges from hidden and value widths rather than the key width. Divide CP attention extras by TP like the retained widths, and say CP above 2 reuses the CP2 allowance. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 34 ++++++++++------- .../test_trainer_rank_checkpoint_memory.py | 37 ++++++++++++++++++- 2 files changed, 57 insertions(+), 14 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index fe45a9ba9..b0a680fbe 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -8020,18 +8020,24 @@ def _recomputed_mixer_bytes_per_token(self) -> int: Full recompute replays one layer with gradients, and its attention or GDN mixer keeps what its backward needs across that layer's MoE stage. - Price the larger mixer the model has, from Qwen3.6-35B-A3B allocator - traces on H200 (bytes per local token, at the layer's recompute peak): - - - attention: 66 KB at CP1, the retained width below. CP ranks also - keep stage-padded Q/K/V, the stage output and a core-attention copy: - 94 KB at CP2, priced at 95 KB. - - GDN: 75 KB at CP1, fewer tensors than the retained width, priced at - 74 KB. CP ranks add rank-exchange copies: 84 KB at CP2, priced at - 86 KB. + Price the larger mixer the model has. Sites and sizes come from + Qwen3.6-35B-A3B allocator traces on H200, per local token at the + layer's recompute peak: + + - attention: the retained attention width above (66 KB measured at + CP1, 69 KB priced). A context-parallel rank also keeps its + stage-padded Q/K/V, the stage output and a core-attention copy + (94 KB measured at CP2, 95 KB priced). CP above 2 uses the CP2 + allowance; ranks with several remote stages may keep more. + - GDN: norm output, q and k (their fp32 l2norm copies count twice), + v, z, two segment-layout tensors, the gated-norm output and the + chunk decay matrix (75 KB measured at CP1, 74 KB priced). A + context-parallel rank adds its hidden-width input and value-width + output exchanges (84 KB measured at CP2, 86 KB priced). """ geometry = self._geometry hidden = self._hidden_size + tp = max(1, self._topology_key()[1]) cp = self._topology_key()[2] > 1 widths = [] if self._gdn_layers < self._num_layers: @@ -8039,14 +8045,16 @@ def _recomputed_mixer_bytes_per_token(self) -> int: if cp: q = geometry.num_attention_heads * geometry.kv_channels or hidden kv = geometry.num_query_groups * geometry.kv_channels or hidden - attention += 3 * q + 2 * kv + attention += (3 * q + 2 * kv) / tp widths.append(attention) if self._gdn_layers: key = geometry.gdn_key_heads * geometry.gdn_key_head_dim value = geometry.gdn_value_heads * geometry.gdn_value_head_dim - widths.append( - 2 * hidden + 4 * key + 6 * value + ((key + value) if cp else 0) - ) + chunk = 64 * geometry.gdn_value_heads + gdn = hidden + (6 * key + 5 * value + chunk) / tp + if cp: + gdn += hidden + value / tp + widths.append(gdn) return int(max(widths, default=0) * self._param_dtype_size) def _gdn_segment_layer_bytes(self) -> float: diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index 8542f5ec0..887994a29 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -517,7 +517,8 @@ def qwen36_attention(r): @pytest.mark.parametrize( "cp,expected", - # Traces measured 66 KB per token at CP1 and 94 KB at CP2. + # Traces measured 66 KB per token at CP1 and 94 KB at CP2. CP4 reuses the + # CP2 stage allowance; it is not measured. [(1, 2 * (2 * 2048 + 7 * 4096 + 3 * 512)), (2, 95232), (4, 95232)], ) def test_recomputed_attention_is_priced_beside_the_moe_stage(cp, expected): @@ -557,3 +558,37 @@ def test_no_grad_groups_do_not_recompute_a_mixer(): r = qwen36_attention(rank()) moe = r._moe_workspace_bytes(10) assert r._checkpoint_memory_floor(((10, False),)) == (0, moe + 4 * 10 * 2048 * 2) + + +def test_ungated_attention_prices_fewer_projections(): + r = qwen36_attention(rank()) + r._attention_output_gate = False + assert r._recomputed_mixer_bytes_per_token() == 2 * (2 * 2048 + 5 * 4096 + 3 * 512) + r._topology_key = lambda: (1, 1, 2, 1) + assert r._recomputed_mixer_bytes_per_token() == 2 * ( + 2 * 2048 + 5 * 4096 + 3 * 512 + 3 * 4096 + 2 * 512 + ) + + +@pytest.mark.parametrize("cp", [1, 2]) +def test_gdn_width_follows_hidden_key_and_value_separately(cp): + # Hidden differs from the key width, as in Qwen3.5-27B: CP exchanges + # carry hidden-width inputs and value-width outputs. + r = rank() + r._hidden_size = r.runtime.model[0].decoder.config.hidden_size = 5120 + r._geometry = replace( + r._geometry, + gdn_key_heads=16, + gdn_key_head_dim=128, + gdn_value_heads=48, + gdn_value_head_dim=128, + ) + r._gdn_layers = r._num_layers + r._topology_key = lambda: (1, 1, cp, 1) + key, value = 16 * 128, 48 * 128 + width = 5120 + 6 * key + 5 * value + 64 * 48 + if cp > 1: + width += 5120 + value + assert r._recomputed_mixer_bytes_per_token() == 2 * width + moe = r._moe_workspace_bytes(7, checkpoint_grad=True) + assert r._checkpoint_memory_floor(((7, True),))[1] == moe + 7 * 2 * width From 91c98a816495ae9509c6b7775a0a4cdb0a2435e8 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 07:20:03 +0000 Subject: [PATCH 04/40] Price GDN l2norm outputs at value-head width The l2-normalized q and k are expanded to the value heads before they are saved, so their width follows value_heads * key_head_dim, not twice the key width. Qwen3.6 is unchanged; geometries with more value than key heads were under-priced. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 14 ++++++++------ tests/unit/test_trainer_rank_checkpoint_memory.py | 5 +++-- 2 files changed, 11 insertions(+), 8 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index b0a680fbe..1e1deb4e0 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -8029,11 +8029,12 @@ def _recomputed_mixer_bytes_per_token(self) -> int: stage-padded Q/K/V, the stage output and a core-attention copy (94 KB measured at CP2, 95 KB priced). CP above 2 uses the CP2 allowance; ranks with several remote stages may keep more. - - GDN: norm output, q and k (their fp32 l2norm copies count twice), - v, z, two segment-layout tensors, the gated-norm output and the - chunk decay matrix (75 KB measured at CP1, 74 KB priced). A - context-parallel rank adds its hidden-width input and value-width - output exchanges (84 KB measured at CP2, 86 KB priced). + - GDN: norm output, q and k, their l2norm outputs expanded to the + value heads, v, z, two segment-layout tensors, the gated-norm + output and the chunk decay matrix (75 KB measured at CP1, 74 KB + priced). A context-parallel rank adds its hidden-width input + exchange (measured) and value-width output projection (from + source): 84 KB measured at CP2, 86 KB priced. """ geometry = self._geometry hidden = self._hidden_size @@ -8050,8 +8051,9 @@ def _recomputed_mixer_bytes_per_token(self) -> int: if self._gdn_layers: key = geometry.gdn_key_heads * geometry.gdn_key_head_dim value = geometry.gdn_value_heads * geometry.gdn_value_head_dim + normalized = 2 * geometry.gdn_value_heads * geometry.gdn_key_head_dim chunk = 64 * geometry.gdn_value_heads - gdn = hidden + (6 * key + 5 * value + chunk) / tp + gdn = hidden + (2 * key + normalized + 5 * value + chunk) / tp if cp: gdn += hidden + value / tp widths.append(gdn) diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index 887994a29..66a8f0e37 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -535,7 +535,7 @@ def test_recomputed_attention_is_priced_beside_the_moe_stage(cp, expected): # Hybrid: GDN (74 KB) is larger at CP1, CP attention (95 KB) at CP2. (1, 30, 73728), (2, 30, 95232), - # GDN only: its CP rank-exchange copies add a key and a value row. + # GDN only: CP exchanges add a hidden-width and a value-width row. (1, 40, 73728), (2, 40, 86016), ], @@ -586,7 +586,8 @@ def test_gdn_width_follows_hidden_key_and_value_separately(cp): r._gdn_layers = r._num_layers r._topology_key = lambda: (1, 1, cp, 1) key, value = 16 * 128, 48 * 128 - width = 5120 + 6 * key + 5 * value + 64 * 48 + # q and k are l2-normalized after expansion to the 48 value heads. + width = 5120 + 2 * key + 2 * 48 * 128 + 5 * value + 64 * 48 if cp > 1: width += 5120 + value assert r._recomputed_mixer_bytes_per_token() == 2 * width From cf56c618812531afd5d1037faa92444de8a402bf Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 15:22:52 +0000 Subject: [PATCH 05/40] Price one dispatched H-wide input per routed row under HybridEP The EP1 all-to-all holds its permuted copy and the exchanged rows at the expert stage; HybridEP permutes while it dispatches and returns one tensor. A Qwen3.6 CP2/EP2 allocator trace holds exactly one routed H-wide input beside the FC1 and FC2 stage tensors (9,728 features per routed row), where the planner charged two (11,776). Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 20 ++++++----- .../test_trainer_rank_converted_memory.py | 36 +++++++++++++++++-- tests/unit/test_trainer_rank_moe_memory.py | 8 +++-- tests/unit/test_trainer_rank_shared_memory.py | 7 ++-- 4 files changed, 54 insertions(+), 17 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 1e1deb4e0..3d2b46fa8 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1606,6 +1606,10 @@ def _moe_output_bytes_per_token( # permuted. Balanced routing gives local tokens x top-k, as at EP1; a # pretrained CP2/EP2 run put about 1.35x that on one rank. routed_allowance = _EP_ROUTED_ROW_ALLOWANCE if shape.ep > 1 else 1 + # Routed H-wide inputs held at the expert stage. The EP1 all-to-all keeps + # its permuted copy and the exchanged rows; HybridEP permutes while it + # dispatches and returns one tensor. + dispatched = 1 if shape.ep > 1 else 2 coefficient = 0 for chunk in model: for layer in chunk.modules(): @@ -1699,8 +1703,6 @@ def _moe_output_bytes_per_token( and fc1.fused_gate_up and not fc1.non_gated and fc1.out_features == 2 * inputs.shape[-2] - # HybridEP keeps one dispatched H-wide input where the - # EP1 all-to-all keeps two; charging two is conservative. and getattr(dispatcher, "ep_size", None) == shape.ep and getattr(dispatcher, "tp_size", None) == 1 and getattr(dispatcher, "num_local_experts", 0) > 1 @@ -1709,10 +1711,10 @@ def _moe_output_bytes_per_token( and getattr(experts, "offload_moe_act", None) is False and getattr(experts, "activation_recompute", None) is False ): - # The two dispatched H-wide inputs and FC1 gate/up sum - # remain live at the FC2 sum, including in the observed - # compiled path. This is one stage, not a backward bound. - features += 2 * fc2.out_features + fc1.out_features + # The dispatched H-wide inputs and FC1 gate/up sum remain + # live at the FC2 sum, including in the observed compiled + # path. This is one stage, not a backward bound. + features += dispatched * fc2.out_features + fc1.out_features enclosing_fc1 = fc1 shared = _shared_expert_output_bytes_per_token(layer) if ( @@ -1756,13 +1758,13 @@ def _moe_output_bytes_per_token( and first_tensors[1].shape[2] == enclosing_fc1.out_features ): first_padding, first_transposes, first_rank = first - # FC1 retains both routed H inputs and its base O1 + # FC1 retains the routed H inputs and its base O1 # while producing adapter O1. Its sum is not live yet. converted_stages.append( ( routed_size * ( - 2 * fc2.out_features + dispatched * fc2.out_features + 2 * enclosing_fc1.out_features + first_rank ) @@ -1776,7 +1778,7 @@ def _moe_output_bytes_per_token( ( routed_size * ( - 2 * fc2.out_features + dispatched * fc2.out_features + 3 * enclosing_fc1.out_features + (first_rank if checkpoint_grad else 0) ) diff --git a/tests/unit/test_trainer_rank_converted_memory.py b/tests/unit/test_trainer_rank_converted_memory.py index 6e8adf9bb..c0c909931 100644 --- a/tests/unit/test_trainer_rank_converted_memory.py +++ b/tests/unit/test_trainer_rank_converted_memory.py @@ -4,13 +4,17 @@ from typing import Any, cast import pytest -from test_trainer_rank_moe_memory import _enclosing_moe, _rank +from test_trainer_rank_moe_memory import _enclosing_moe, _hybridep, _rank from test_trainer_rank_moe_memory import layer as layer from test_trainer_rank_pending_memory import module, rank_with_moe import torch from art.trainer_rank import ForwardInput, _gdn_memory -from art.trainer_rank._impl import _expert_lora_weight_storage +from art.trainer_rank._impl import ( + _expert_lora_weight_storage, + _moe_output_bytes_per_token, +) +from art.trainer_rank._planner_cost import ParallelShape def weights(layer: Any, rank: int, *, fc1: bool = True, dtype=torch.bfloat16): @@ -266,6 +270,34 @@ def test_fc1_fixed_weights_do_not_scale_with_topk(layer, topk): ) +def test_hybridep_fc1_stages_hold_one_dispatched_input(layer): + # FC1's converted stages hold the routed H-wide inputs too: two under the + # EP1 all-to-all, one under HybridEP, over its 12 allowance rows. + weights(layer, 8) + single: list[tuple[int, int]] = [] + _moe_output_bytes_per_token( + [layer], ParallelShape(tp=1, cp=1), checkpoint_grad=True, converted_stages=single + ) + expert = _hybridep(layer, 2) + expert.token_dispatcher.num_local_experts = 128 + sharded: list[tuple[int, int]] = [] + _moe_output_bytes_per_token( + [expert], + ParallelShape(tp=1, cp=2, ep=2), + checkpoint_grad=True, + converted_stages=sharded, + ) + assert [stage[0] for stage in single[:2]] == [ + 16 * (2 * 2048 + 2 * 1024 + 8), + 16 * (2 * 2048 + 3 * 1024 + 8), + ] + assert [stage[0] for stage in sharded[:2]] == [ + 24 * (2048 + 2 * 1024 + 8), + 24 * (2048 + 3 * 1024 + 8), + ] + assert [stage[1] for stage in sharded] == [stage[1] for stage in single] + + def test_wide_fc1_sum_is_a_separate_stage(layer): weights(layer, 8) experts = layer.experts diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index aa30b54d8..66a114424 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -215,8 +215,10 @@ def test_hybridep_prices_routed_rows_with_imbalance_allowance(layer, ep): def test_hybridep_keeps_the_enclosing_fc1_stage(layer): - # The FC1 inputs and gate/up sum stay live at the FC2 sum under HybridEP - # too; EP1's two dispatched H-wide inputs over-count HybridEP's one. + # The FC1 input and gate/up sum stay live at the FC2 sum under HybridEP + # too. The EP1 all-to-all holds two routed H-wide inputs (its permuted + # copy and the exchanged rows); HybridEP permutes while dispatching and + # holds one, as a Qwen3.6 CP2/EP2 allocator trace shows. single = _moe_output_bytes_per_token( [_enclosing_moe(layer)], ParallelShape(tp=1, cp=1) ) @@ -224,7 +226,7 @@ def test_hybridep_keeps_the_enclosing_fc1_stage(layer): expert.token_dispatcher.num_local_experts = 128 sharded = _moe_output_bytes_per_token([expert], ParallelShape(tp=1, cp=2, ep=2)) assert single == 8 * (512 + 3 * 2048 + 2 * 2048 + 1024) * 2 - assert sharded == single // 8 * 12 + assert sharded == 12 * (512 + 3 * 2048 + 2048 + 1024) * 2 @pytest.mark.parametrize( diff --git a/tests/unit/test_trainer_rank_shared_memory.py b/tests/unit/test_trainer_rank_shared_memory.py index 36bc971d7..6867e3566 100644 --- a/tests/unit/test_trainer_rank_shared_memory.py +++ b/tests/unit/test_trainer_rank_shared_memory.py @@ -410,11 +410,12 @@ def test_shared_return_escapes_the_ep_routed_allowance(layer, checkpoint_grad, s layer.config.context_parallel_size = layer.config.expert_model_parallel_size = 2 # EP2 halves the experts each rank owns. _hybridep(layer, 2).token_dispatcher.num_local_experts = local_experts // 2 - # HybridEP's 1.5x allowance turns top-k 8 into 12 routed rows; the gated - # shared return (doubled for checkpoint backward) is per local token. + # HybridEP's 1.5x allowance turns top-k 8 into 12 routed rows, each with + # one dispatched H-wide input; the gated shared return (doubled for + # checkpoint backward) is per local token. assert ( _moe_output_bytes_per_token( [layer], ParallelShape(tp=1, cp=2, ep=2), checkpoint_grad=checkpoint_grad ) - == 12 * 11776 * 2 + shared + == 12 * 9728 * 2 + shared ) From d583cd210ef4d037f6b4f37d5c26c965682e548a Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 15:54:48 +0000 Subject: [PATCH 06/40] Format the HybridEP converted-stage test Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_trainer_rank_converted_memory.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_trainer_rank_converted_memory.py b/tests/unit/test_trainer_rank_converted_memory.py index c0c909931..5bba9ac56 100644 --- a/tests/unit/test_trainer_rank_converted_memory.py +++ b/tests/unit/test_trainer_rank_converted_memory.py @@ -276,7 +276,10 @@ def test_hybridep_fc1_stages_hold_one_dispatched_input(layer): weights(layer, 8) single: list[tuple[int, int]] = [] _moe_output_bytes_per_token( - [layer], ParallelShape(tp=1, cp=1), checkpoint_grad=True, converted_stages=single + [layer], + ParallelShape(tp=1, cp=1), + checkpoint_grad=True, + converted_stages=single, ) expert = _hybridep(layer, 2) expert.token_dispatcher.num_local_experts = 128 From 210acec08628365264bf1dfffd6df895cd7312bc Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 15:58:45 +0000 Subject: [PATCH 07/40] Name the EP1 all-to-all's second routed copy as its expert-sorted rows Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 6 +++--- tests/unit/test_trainer_rank_moe_memory.py | 6 +++--- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 3d2b46fa8..c14a9254b 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1606,9 +1606,9 @@ def _moe_output_bytes_per_token( # permuted. Balanced routing gives local tokens x top-k, as at EP1; a # pretrained CP2/EP2 run put about 1.35x that on one rank. routed_allowance = _EP_ROUTED_ROW_ALLOWANCE if shape.ep > 1 else 1 - # Routed H-wide inputs held at the expert stage. The EP1 all-to-all keeps - # its permuted copy and the exchanged rows; HybridEP permutes while it - # dispatches and returns one tensor. + # Routed H-wide inputs held at the expert stage. The EP1 all-to-all path + # keeps its permuted rows and their expert-sorted copy; HybridEP permutes + # while it dispatches and returns one tensor. dispatched = 1 if shape.ep > 1 else 2 coefficient = 0 for chunk in model: diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 66a114424..d1884067d 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -216,9 +216,9 @@ def test_hybridep_prices_routed_rows_with_imbalance_allowance(layer, ep): def test_hybridep_keeps_the_enclosing_fc1_stage(layer): # The FC1 input and gate/up sum stay live at the FC2 sum under HybridEP - # too. The EP1 all-to-all holds two routed H-wide inputs (its permuted - # copy and the exchanged rows); HybridEP permutes while dispatching and - # holds one, as a Qwen3.6 CP2/EP2 allocator trace shows. + # too. The EP1 all-to-all path holds two routed H-wide inputs (its + # permuted rows and their expert-sorted copy); HybridEP permutes while + # dispatching and holds one, as Qwen3.6 CP2 allocator traces show. single = _moe_output_bytes_per_token( [_enclosing_moe(layer)], ParallelShape(tp=1, cp=1) ) From c59a692fe6f536ca877b2768eb81a62598ec31dc Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 18:50:52 +0000 Subject: [PATCH 08/40] Price what is live at the recompute peak; make the EP allowance EP-dependent Backward recomputes the last layer first, so the checkpoint floor's peak meets every saved boundary but only the one incoming gradient. Where the MoE stage is priced, charge that gradient instead of one per boundary (39 hidden rows per token too many at 40 layers), and price what the old allowance was silently covering, all from Qwen3.6-35B-A3B allocator traces: - the recomputed layer's residual and pre-MLP norm output (2H per row); - GDN's sixth value-width tensor (the projected q/k/v includes v); - the shared expert's saved FC1 gate/up and GLU outputs; - router scores and map plus the dispatcher's row-id map (EP1) or probability copy and handle (HybridEP); - TE's cuBLAS workspaces, as growth until its GEMMs allocate them. Without a priced MoE stage the per-boundary allowance stays: it also covers dense MLP and other recompute work the floor does not price. The EP>1 routed-row allowance becomes EP-dependent (1.4, 1.6, 2.0 at EP2, 4, 8), from pretrained Qwen3.6 routing of 3.5M tokens of retail agent trajectories (worst layer 1.21, 1.41, 1.63) and one production EP2 run (1.35). Routed rows are no longer rounded up to whole rows. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 140 +++++++++++++++--- ...trainer_rank_checkpoint_gradient_memory.py | 42 ++++-- .../test_trainer_rank_checkpoint_memory.py | 70 +++++++-- .../test_trainer_rank_converted_memory.py | 33 +++-- tests/unit/test_trainer_rank_head_memory.py | 6 +- tests/unit/test_trainer_rank_moe_memory.py | 25 +++- .../unit/test_trainer_rank_pending_memory.py | 8 +- tests/unit/test_trainer_rank_shared_memory.py | 24 +-- 8 files changed, 267 insertions(+), 81 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index c14a9254b..222b6a371 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1532,8 +1532,29 @@ def _expert_lora_weight_storage( return (transposes if a.shape[2] < 8 else 0, transposes, effective) -# Routed rows per rank at EP>1, relative to balanced routing (see below). -_EP_ROUTED_ROW_ALLOWANCE = 1.5 +# Routed rows on the most loaded rank at EP>1, relative to balanced routing. +# Expert-shard load is uneven per layer, from the router's expert preferences, +# and larger batches do not average it away. Pretrained Qwen3.6-35B-A3B on 3.5M +# tokens of retail agent trajectories: worst layer 1.21 at EP2, 1.41 at EP4, +# 1.63 at EP8 (1.88 on a small rollout sample); one production EP2 run saw +# 1.35. These samples bound what was measured, not all routing. Unmeasured EP +# sizes use the next measured one; above EP8 the allowance grows with log2(EP) +# up to EP itself (every pair on one rank). +_EP_ROUTED_ROW_ALLOWANCE = {2: 1.4, 4: 1.6, 8: 2.0} + + +def _ep_routed_row_allowance(ep: int) -> float: + if ep <= 1: + return 1.0 + for size, allowance in sorted(_EP_ROUTED_ROW_ALLOWANCE.items()): + if ep <= size: + return allowance + return min(float(ep), 2.0 + 0.4 * math.log2(ep / 8)) + + +# Transformer Engine's Hopper cuBLAS workspaces: one per grouped-GEMM stream +# (four) plus the plain GEMM's, each 32 MiB + 1 KiB. +_TE_CUBLAS_WORKSPACE_BYTES = 5 * (32 * 2**20 + 1024) def _moe_dispatcher_supported( @@ -1603,9 +1624,8 @@ def _moe_output_bytes_per_token( from art.megatron.lora import LoRA, MLPExpertsLinearFC1LoRA, MLPExpertsLinearFC2LoRA # HybridEP hands each rank the pairs routed to its local experts, already - # permuted. Balanced routing gives local tokens x top-k, as at EP1; a - # pretrained CP2/EP2 run put about 1.35x that on one rank. - routed_allowance = _EP_ROUTED_ROW_ALLOWANCE if shape.ep > 1 else 1 + # permuted. Balanced routing gives local tokens x top-k, as at EP1. + routed_allowance = _ep_routed_row_allowance(shape.ep) # Routed H-wide inputs held at the expert stage. The EP1 all-to-all path # keeps its permuted rows and their expert-sorted copy; HybridEP permutes # while it dispatches and returns one tensor. @@ -1726,8 +1746,10 @@ def _moe_output_bytes_per_token( # Gate-score backward saves a distinct pre-gate X. Charge it # beside this layer's returned X, not another layer's maximum. shared += shared - routed_rows = math.ceil(config.moe_router_topk * routed_allowance) - row_bytes = routed_rows * features * weights.element_size() + shared + routed_rows = config.moe_router_topk * routed_allowance + row_bytes = ( + math.ceil(routed_rows * features * weights.element_size()) + shared + ) coefficient = max(coefficient, row_bytes) storage = _expert_lora_weight_storage(lora, slot_ref) if converted_stages is not None and storage is not None: @@ -1832,6 +1854,11 @@ def _moe_output_bytes_per_token( + 2 * (experts_count + 1) * 4, ) ) + if converted_stages is not None: + # The EP allowance gives fractional routed rows; round each stage up. + converted_stages[:] = [ + (math.ceil(per_row), fixed) for per_row, fixed in converted_stages + ] return coefficient @@ -4093,17 +4120,81 @@ def _checkpoint_memory_floor( ): return 0, 0 retained = gradient_rows * layers * self._hidden_size * 2 - if gradient_rows: - self._checkpoint_moe_bytes_per_token() + moe = self._checkpoint_moe_bytes_per_token() if gradient_rows else 0 refs = (None,) * len(group_rows) if slot_refs is None else slot_refs - mixer = self._recomputed_mixer_bytes_per_token() if gradient_rows else 0 + # Beside the mixer, the recomputed layer keeps its post-mixer residual + # and pre-MLP norm output, and its MoE stage its routing state. + mixer = ( + self._recomputed_mixer_bytes_per_token() + + ( + 2 * self._hidden_size * 2 + self._moe_checkpoint_state_bytes_per_token() + if moe + else 0 + ) + if gradient_rows + else 0 + ) workspace = max( self._moe_workspace_bytes(rows, checkpoint_grad=grad, slot_ref=ref) + (mixer * rows if grad else 4 * rows * self._hidden_size * 2) for (rows, grad), ref in zip(group_rows, refs, strict=True) ) + if moe: + workspace += self._te_workspace_growth_bytes() return retained, workspace + def _te_workspace_growth_bytes(self) -> int: + """Transformer Engine's cuBLAS workspaces, until its GEMMs allocate them. + + The first plain and grouped GEMMs allocate them during a call, and TE + keeps them for the process; later calls see them as used memory. + """ + try: + from transformer_engine.pytorch.cpp_extensions import gemm + except ImportError: + return _TE_CUBLAS_WORKSPACE_BYTES + info = getattr(gemm.get_cublas_workspace, "cache_info", None) + if callable(info) and info().currsize >= 2: + return 0 + return _TE_CUBLAS_WORKSPACE_BYTES + + def _moe_checkpoint_state_bytes_per_token(self) -> int: + """Per local token beside the recomputed MoE stage's routed rows. + + FP32 router scores and the boolean routing map; the dispatcher's state + (the EP1 permutation's int32 row-id map of 2E + 1, or HybridEP's FP32 + probability copy and handle metadata of about 5E); and the shared + expert's saved FC1 gate/up and GLU outputs. Qwen3.6-35B-A3B traces: + 3,332 and 3,593 bytes of routing state at EP1 and EP2, 3 KB shared. + """ + geometry = self._geometry + experts = geometry.moe_experts + if not experts: + return 0 + routing = experts * (4 + 1) + ( + 4 * (2 * experts + 1) if self._parallel_shape.ep == 1 else 9 * experts + 16 + ) + return routing + 3 * geometry.moe_shared_expert_ffn * self._param_dtype_size + + def _checkpoint_input_gradient_bytes( + self, group_rows: tuple[tuple[int, bool], ...] + ) -> int: + """Gradient rows live at the recomputed layer's peak. + + Backward recomputes the last layer first, so its peak meets every saved + boundary but only the one incoming gradient. Where the MoE stage is + priced (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. + """ + retained, _ = self._checkpoint_memory_floor(group_rows) + if not retained: + return 0 + gradient_rows = sum(rows for rows, grad in group_rows if grad) + if self._checkpoint_moe_bytes_per_token(): + return gradient_rows * self._hidden_size * 2 + return retained + def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( packed_tokens=plan.packed_tokens, @@ -4161,11 +4252,9 @@ def _subforward_cost( checkpoint_floor[0], ), ) - # One logical BF16 input gradient per eligible full/uniform/1 boundary. - # This partial peak allowance is not evidence of simultaneous distinct - # backing stores, nor a bound for compiler saves or other backward work. - # Keep it out of forward retention, including the cold fallback above. - gradient = checkpoint_retained + # Input gradients live at the recomputed layer's peak; kept out of + # forward retention, including the cold fallback above. + gradient = self._checkpoint_input_gradient_bytes(group_rows) checkpoint_retained = output_bytes + max( checkpoint_retained, checkpoint_floor[0] ) @@ -7910,7 +7999,11 @@ def _estimate_required_memory_bytes_from_values( static_compute, max(retained, checkpoint_floor[0]) + max(workspace, head_workspace_bytes, checkpoint_floor[1]) - + (retained if include_checkpoint_input_gradient else 0), + + ( + self._checkpoint_input_gradient_bytes(group_rows) + if include_checkpoint_input_gradient + else 0 + ), ) if signature.topology[2] > 1: # Local head results coexist with full CP outputs during gathering. @@ -8031,12 +8124,13 @@ def _recomputed_mixer_bytes_per_token(self) -> int: stage-padded Q/K/V, the stage output and a core-attention copy (94 KB measured at CP2, 95 KB priced). CP above 2 uses the CP2 allowance; ranks with several remote stages may keep more. - - GDN: norm output, q and k, their l2norm outputs expanded to the - value heads, v, z, two segment-layout tensors, the gated-norm - output and the chunk decay matrix (75 KB measured at CP1, 74 KB - priced). A context-parallel rank adds its hidden-width input - exchange (measured) and value-width output projection (from - source): 84 KB measured at CP2, 86 KB priced. + - GDN: norm output, the projected q/k/v, their l2norm outputs + expanded to the value heads, five more value-width tensors (z, two + segment-layout tensors, the gated-norm output and its gated + product) and the chunk decay matrix: 80 KB measured at CP1, 82 KB + priced. A context-parallel rank's exchanged layout holds about 7% + more rows; its hidden-width input exchange and value-width output + allowance price that (88 KB measured at CP2, 94 KB priced). """ geometry = self._geometry hidden = self._hidden_size @@ -8055,7 +8149,7 @@ def _recomputed_mixer_bytes_per_token(self) -> int: value = geometry.gdn_value_heads * geometry.gdn_value_head_dim normalized = 2 * geometry.gdn_value_heads * geometry.gdn_key_head_dim chunk = 64 * geometry.gdn_value_heads - gdn = hidden + (2 * key + normalized + 5 * value + chunk) / tp + gdn = hidden + (2 * key + normalized + 6 * value + chunk) / tp if cp: gdn += hidden + value / tp widths.append(gdn) diff --git a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py index d6d9ded56..b4008f6c3 100644 --- a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py @@ -1,4 +1,9 @@ -"""Partial input-gradient extents: CPU admission math, not peak/overlap proof.""" +"""Input-gradient extents: CPU admission math, not peak/overlap proof. + +Where the recomputed MoE stage is priced, the floor charges the one incoming +gradient live at the last layer's recompute peak; elsewhere it keeps one +gradient per saved boundary. +""" from dataclasses import replace @@ -25,25 +30,33 @@ def test_pending_cold_peak_does_not_become_forward_retention(pending_rank): r = pending_rank plan = r._plan_flat_forward(full_requests()) cost = r._plan_cost(plan) - gradient = 8 * 6330 * 40 * 2048 * 2 + boundaries = 8 * 6330 * 40 * 2048 * 2 + gradient = 8 * 6330 * 2048 * 2 assert cost.checkpoint_input_gradient == gradient - # Exact previous cold estimate, including outputs and its one safety factor. + # Exact cold estimate, including outputs and its one safety factor. assert cost.retained == 23102959299 - assert cost.required == int((plan.output_bytes + 2 * gradient + 12705630112) * 1.1) + assert cost.required == int( + (plan.output_bytes + boundaries + gradient + cost.checkpoint_workspace) * 1.1 + ) assert r._memory_check(plan).estimated_required_bytes == cost.required profile(r, plan) warm = r._plan_cost(plan) - assert warm.retained == int((plan.output_bytes + gradient) * 1.1) + assert warm.retained == int((plan.output_bytes + boundaries) * 1.1) assert warm.required == cost.required +@pytest.mark.parametrize("moe", [True, False]) @pytest.mark.parametrize("rows", [1, 67, 1024]) -def test_attention_only_extent_scales_with_gradient_rows(rows): +def test_attention_only_extent_scales_with_gradient_rows(rows, moe): r = rank() + if not moe: + # Without a priced MoE stage, one gradient per boundary stays: it also + # covers the dense MLP and other recompute work the floor omits. + r._moe_output_bytes_per_token = r._moe_checkpoint_grad_bytes_per_token = 0 values = r._estimate_flat_forward(requests(rows, 4096)) cost = price(r, values) assert r._gdn_layers == 0 - assert cost.checkpoint_input_gradient == rows * 40 * 2048 * 2 + assert cost.checkpoint_input_gradient == rows * (40 if not moe else 1) * 2048 * 2 assert cost.required >= int( ( cost.checkpoint_retained @@ -59,10 +72,11 @@ def test_gradient_is_not_absorbed_by_larger_head_workspace(): n, out, sig, groups, _head = r._estimate_flat_forward(requests(67, 4096)) head = 10**10 cost = price(r, (n, out, sig, groups, head)) - gradient = 67 * 40 * 2048 * 2 + boundaries = 67 * 40 * 2048 * 2 + gradient = 67 * 2048 * 2 assert cost.checkpoint_workspace == head - assert cost.required == int((out + head + 2 * gradient) * 1.1) - assert cost.retained == int((out + head + gradient) * 1.1) + assert cost.required == int((out + head + boundaries + gradient) * 1.1) + assert cost.retained == int((out + head + boundaries) * 1.1) r._memory_profiles[sig] = _MemoryProfile( bytes_per_token=10**9, packed_tokens=n, @@ -111,11 +125,7 @@ def test_split_sums_all_gradient_children_outside_workspace_max(): profile(r, child) costs = [r._plan_cost(child) for child in children] split = _SplitForwardPlan(tuple(children), ((0,), (1,), (2,)), 3) - assert [c.checkpoint_input_gradient for c in costs] == [ - 17 * 40 * 4096, - 29 * 40 * 4096, - 0, - ] + assert [c.checkpoint_input_gradient for c in costs] == [17 * 4096, 29 * 4096, 0] expected = int( ( sum(c.checkpoint_retained + c.checkpoint_input_gradient for c in costs) @@ -157,7 +167,7 @@ def test_lower_bound_profile_cliff_preserves_separate_peak_component(): lower = r._split_chunk_lower_cost( req, tuple(q.input_tokens for q in req), checkpoint=Unset ) - assert lower.checkpoint_input_gradient == 128 * 40 * 4096 + assert lower.checkpoint_input_gradient == 128 * 4096 assert lower.required <= r._plan_cost(full).required assert lower.checkpoint_retained == full.output_bytes + 128 * 40 * 4096 diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index 66a8f0e37..f73a46e73 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -8,7 +8,12 @@ import torch from art.trainer_rank import ForwardInput, TrainerRank -from art.trainer_rank._impl import Unset, _ForwardRefusal, _MemoryProfile +from art.trainer_rank._impl import ( + _TE_CUBLAS_WORKSPACE_BYTES, + Unset, + _ForwardRefusal, + _MemoryProfile, +) def rank(): @@ -105,7 +110,9 @@ def test_required_and_learned_retained_use_max_not_sum(): ) cost = price(r, values) old = max(n * 2048 * 2 * 14, n * 188416) - assert cost.required == int((out + max(old, 2 * retained + work)) * 1.1) + gradient = r._checkpoint_input_gradient_bytes(groups) + assert gradient == 1024 * 2048 * 2 # The incoming gradient only. + assert cost.required == int((out + max(old, retained + gradient + work)) * 1.1) assert cost.retained == int((out + retained) * 1.1) r._memory_profiles[sig] = replace( r._memory_profiles[sig], @@ -135,7 +142,7 @@ def test_per_group_padding_precedes_gradient_filter(): assert values[0] == 40 and values[3] == ((16, True), (24, False)) assert r._checkpoint_memory_floor(values[3]) == ( 16 * 40 * 4096, - 24 * (188416 + 4 * 2048 * 2), + 24 * (188416 + 4 * 2048 * 2) + _TE_CUBLAS_WORKSPACE_BYTES, ) @@ -325,7 +332,8 @@ def test_split_keeps_complete_order_and_checks_each_new_subforward(): logical_per_packed=1, retained_compute_bytes_per_token=1, ) - limit = 250_000_000 + # Just below the unsplit plan's requirement: two subforwards must fit. + limit = r._plan_cost(flat).required - 1 used = 0 r._available_memory_bytes = lambda: limit - used result = r._find_admissible_forward(req, checkpoint=Unset, refusal_prefix="test") @@ -431,7 +439,8 @@ def test_no_grad_enclosure_uses_max_group_and_affine_stage(): mixed = ((3, True), (11, False)) assert r._checkpoint_memory_floor(mixed) == ( 3 * 40 * 2048 * 2, - max(3 * 188416, max(11 * 188416, 11 + 1_000_000) + 4 * 11 * 2048 * 2), + max(3 * 188416, max(11 * 188416, 11 + 1_000_000) + 4 * 11 * 2048 * 2) + + _TE_CUBLAS_WORKSPACE_BYTES, ) @@ -526,18 +535,23 @@ def test_recomputed_attention_is_priced_beside_the_moe_stage(cp, expected): r._topology_key = lambda: (1, 1, cp, 1) assert r._recomputed_mixer_bytes_per_token() == expected moe = r._moe_workspace_bytes(10, checkpoint_grad=True) - assert r._checkpoint_memory_floor(((10, True),))[1] == moe + 10 * expected + # Beside the mixer: the recomputed layer's residual and pre-MLP norm rows, + # and Transformer Engine's cuBLAS workspaces. + assert ( + r._checkpoint_memory_floor(((10, True),))[1] + == moe + 10 * (expected + 2 * 2048 * 2) + _TE_CUBLAS_WORKSPACE_BYTES + ) @pytest.mark.parametrize( "cp,gdn_layers,expected", [ - # Hybrid: GDN (74 KB) is larger at CP1, CP attention (95 KB) at CP2. - (1, 30, 73728), + # Hybrid: GDN (82 KB) is larger at CP1, CP attention (95 KB) at CP2. + (1, 30, 81920), (2, 30, 95232), # GDN only: CP exchanges add a hidden-width and a value-width row. - (1, 40, 73728), - (2, 40, 86016), + (1, 40, 81920), + (2, 40, 94208), ], ) def test_larger_recomputed_mixer_is_priced(cp, gdn_layers, expected): @@ -587,9 +601,41 @@ def test_gdn_width_follows_hidden_key_and_value_separately(cp): r._topology_key = lambda: (1, 1, cp, 1) key, value = 16 * 128, 48 * 128 # q and k are l2-normalized after expansion to the 48 value heads. - width = 5120 + 2 * key + 2 * 48 * 128 + 5 * value + 64 * 48 + width = 5120 + 2 * key + 2 * 48 * 128 + 6 * value + 64 * 48 if cp > 1: width += 5120 + value assert r._recomputed_mixer_bytes_per_token() == 2 * width moe = r._moe_workspace_bytes(7, checkpoint_grad=True) - assert r._checkpoint_memory_floor(((7, True),))[1] == moe + 7 * 2 * width + assert ( + r._checkpoint_memory_floor(((7, True),))[1] + == moe + 7 * 2 * (width + 2 * 5120) + _TE_CUBLAS_WORKSPACE_BYTES + ) + + +@pytest.mark.parametrize("ep,routing", [(1, 3332), (2, 3600)]) +def test_moe_state_beside_the_recomputed_stage(ep, routing): + # Qwen3.6-35B-A3B: 256 experts, 512-wide shared expert. Traces measured + # 3,332 (EP1) and 3,593 (EP2) bytes of routing state per local token. + r = rank() + r._geometry = replace(r._geometry, moe_experts=256, moe_shared_expert_ffn=512) + r._parallel_shape = replace(r._parallel_shape, ep=ep) + assert r._moe_checkpoint_state_bytes_per_token() == routing + 3 * 512 * 2 + r._geometry = replace(r._geometry, moe_experts=0) + assert r._moe_checkpoint_state_bytes_per_token() == 0 + + +def test_te_workspaces_are_growth_until_allocated(monkeypatch): + gemm = pytest.importorskip("transformer_engine.pytorch.cpp_extensions.gemm") + r = rank() + entries = [0] + + class Cached: + def cache_info(self): + return SimpleNamespace(currsize=entries[0]) + + monkeypatch.setattr(gemm, "get_cublas_workspace", Cached()) + assert r._te_workspace_growth_bytes() == _TE_CUBLAS_WORKSPACE_BYTES + entries[0] = 1 # Plain GEMM only: the grouped streams are still to come. + assert r._te_workspace_growth_bytes() == _TE_CUBLAS_WORKSPACE_BYTES + entries[0] = 2 + assert r._te_workspace_growth_bytes() == 0 diff --git a/tests/unit/test_trainer_rank_converted_memory.py b/tests/unit/test_trainer_rank_converted_memory.py index 5bba9ac56..fe1b47372 100644 --- a/tests/unit/test_trainer_rank_converted_memory.py +++ b/tests/unit/test_trainer_rank_converted_memory.py @@ -1,5 +1,6 @@ """Source-derived affine routed-expert stages; no complete backward/compiled bound.""" +import math from types import SimpleNamespace from typing import Any, cast @@ -100,12 +101,19 @@ def test_actual_plan_cost_and_admission(layer, rank_value, grad, output): retained, workspace = rank._checkpoint_memory_floor(rank._plan_group_rows(plan)) pending = _gdn_memory.plan_floor(rank, plan) if grad: - # The recomputed layer's mixer stays live beside its MoE stage. + # The recomputed layer's mixer, its residual and pre-MLP norm rows and + # its MoE routing state stay live beside its MoE stage; the first call + # also allocates TE's cuBLAS workspaces. mixer = rank._recomputed_mixer_bytes_per_token() + beside = 2 * 2048 * 2 + rank._moe_checkpoint_state_bytes_per_token() assert mixer > 0 - assert workspace == expected(8, rank_value, True) + 8 * mixer + assert workspace == ( + expected(8, rank_value, True) + + 8 * (mixer + beside) + + rank._te_workspace_growth_bytes() + ) + # The GDN pending floor combines with this one by maximum. assert pending[0] == retained == 8 * 40 * 2048 * 2 - assert pending[1] >= workspace else: assert retained == 0 and pending == (0, 0) assert workspace == expected(8, rank_value, False) + 4 * 8 * 2048 * 2 @@ -131,9 +139,15 @@ def test_reference_and_gradient_keep_distinct_stage_modes(layer, order, rank_val assert set(groups) == {(3, True), (9, False)} retained, workspace = rank._checkpoint_memory_floor(groups) assert retained == 3 * 40 * 2048 * 2 - assert workspace == max( - expected(3, rank_value, True) + 3 * rank._recomputed_mixer_bytes_per_token(), - expected(9, rank_value, False) + 4 * 9 * 2048 * 2, + beside = 2 * 2048 * 2 + rank._moe_checkpoint_state_bytes_per_token() + assert ( + workspace + == max( + expected(3, rank_value, True) + + 3 * (rank._recomputed_mixer_bytes_per_token() + beside), + expected(9, rank_value, False) + 4 * 9 * 2048 * 2, + ) + + rank._te_workspace_growth_bytes() ) assert ( rank._memory_check(plan).estimated_required_bytes @@ -272,7 +286,7 @@ def test_fc1_fixed_weights_do_not_scale_with_topk(layer, topk): def test_hybridep_fc1_stages_hold_one_dispatched_input(layer): # FC1's converted stages hold the routed H-wide inputs too: two under the - # EP1 all-to-all, one under HybridEP, over its 12 allowance rows. + # EP1 all-to-all, one under HybridEP, over its EP2 allowance rows. weights(layer, 8) single: list[tuple[int, int]] = [] _moe_output_bytes_per_token( @@ -294,9 +308,10 @@ def test_hybridep_fc1_stages_hold_one_dispatched_input(layer): 16 * (2 * 2048 + 2 * 1024 + 8), 16 * (2 * 2048 + 3 * 1024 + 8), ] + # 8 x 1.4 routed rows at EP2, rounded up per stage. assert [stage[0] for stage in sharded[:2]] == [ - 24 * (2048 + 2 * 1024 + 8), - 24 * (2048 + 3 * 1024 + 8), + math.ceil(8 * 1.4 * 2 * (2048 + 2 * 1024 + 8)), + math.ceil(8 * 1.4 * 2 * (2048 + 3 * 1024 + 8)), ] assert [stage[1] for stage in sharded] == [stage[1] for stage in single] diff --git a/tests/unit/test_trainer_rank_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py index 21b7ad071..9bfaf9bfc 100644 --- a/tests/unit/test_trainer_rank_head_memory.py +++ b/tests/unit/test_trainer_rank_head_memory.py @@ -143,7 +143,7 @@ def test_outputs_retention_and_empirical_peak_are_counted_once(): r = rank() plan = r._plan_flat_forward([request(grad=True)]) retained = 512 * 40 * 2048 * 2 - gradient = 512 * 40 * 2048 * 2 + gradient = 512 * 2048 * 2 # The incoming gradient beside the MoE stage. head = 3 * 512 * 248320 * 2 cost = r._plan_cost(plan) assert cost.retained == int((plan.output_bytes + retained + head) * 1.1) @@ -308,9 +308,11 @@ def test_tied_standard_head_weight_uses_the_same_capacity(): @pytest.mark.parametrize("rows", [128, 512]) def test_target_backward_refuses_budget_below_logits_and_both_gradients(rows): r = rank() + # Isolate the head term from TE's one-time cuBLAS workspace growth. + r._te_workspace_growth_bytes = lambda: 0 plan = r._plan_flat_forward([request(rows, grad=True)]) retained, _ = r._checkpoint_memory_floor(r._plan_group_rows(plan)) - gradient = rows * 40 * 2048 * 2 + gradient = rows * 2048 * 2 dense = min(rows, 512) * 248320 * 2 before = int((plan.output_bytes + retained + gradient + 2 * dense) * 1.1) expected = int((plan.output_bytes + retained + gradient + 3 * dense) * 1.1) diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index d1884067d..51d8373a5 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -1,6 +1,7 @@ """CPU contracts for one known MoE component, not a whole-model memory bound.""" from dataclasses import replace +import math from types import SimpleNamespace from typing import Any, cast import weakref @@ -11,6 +12,7 @@ from art.trainer_rank import ForwardInput, TrainerRank from art.trainer_rank._impl import ( _PACKED_PRICED_LOGICAL_ROW_BYTES, + _ep_routed_row_allowance, _MemoryProfile, _MemorySignature, _moe_output_bytes_per_token, @@ -203,15 +205,26 @@ def _hybridep(layer, ep: int, manager: str = "hybridep"): return layer -@pytest.mark.parametrize("ep", [2, 4]) -def test_hybridep_prices_routed_rows_with_imbalance_allowance(layer, ep): +@pytest.mark.parametrize("ep,allowance", [(2, 1.4), (4, 1.6), (8, 2.0), (16, 2.4)]) +def test_hybridep_prices_routed_rows_with_imbalance_allowance(layer, ep, allowance): # HybridEP gives each rank the pairs routed to its local experts: balanced - # routing matches EP1's top-k rows per local token, with a 1.5x allowance. + # routing matches EP1's top-k rows per local token, scaled by the measured + # EP-dependent worst-layer imbalance (log2 growth beyond EP8). single = _moe_output_bytes_per_token([layer], ParallelShape(tp=1, cp=1)) sharded = ParallelShape(tp=1, cp=ep, ep=ep) expert = _moe_output_bytes_per_token([_hybridep(layer, ep)], sharded) - # Top-k 8 routed rows per local token become 12; no shared experts here. - assert single > 0 and expert == single // 8 * 12 + # Top-k 8 routed rows per local token become 8 x allowance; no shared + # experts here. + assert _ep_routed_row_allowance(ep) == pytest.approx(allowance) + assert single > 0 and expert == math.ceil(single * allowance) + + +def test_unmeasured_ep_sizes_use_the_next_measured_allowance(): + assert _ep_routed_row_allowance(1) == 1.0 + assert _ep_routed_row_allowance(3) == _ep_routed_row_allowance(4) == 1.6 + assert _ep_routed_row_allowance(6) == 2.0 + # Never above every pair on one rank. + assert _ep_routed_row_allowance(1024) <= 1024 def test_hybridep_keeps_the_enclosing_fc1_stage(layer): @@ -226,7 +239,7 @@ def test_hybridep_keeps_the_enclosing_fc1_stage(layer): expert.token_dispatcher.num_local_experts = 128 sharded = _moe_output_bytes_per_token([expert], ParallelShape(tp=1, cp=2, ep=2)) assert single == 8 * (512 + 3 * 2048 + 2 * 2048 + 1024) * 2 - assert sharded == 12 * (512 + 3 * 2048 + 2048 + 1024) * 2 + assert sharded == math.ceil(8 * 1.4 * (512 + 3 * 2048 + 2048 + 1024) * 2) @pytest.mark.parametrize( diff --git a/tests/unit/test_trainer_rank_pending_memory.py b/tests/unit/test_trainer_rank_pending_memory.py index 7fe3cf5eb..0d53d1f10 100644 --- a/tests/unit/test_trainer_rank_pending_memory.py +++ b/tests/unit/test_trainer_rank_pending_memory.py @@ -133,7 +133,7 @@ def test_actual_constructor_cache_and_full_plan(pending_rank): assert ( rank._memory_check(plan).estimated_required_bytes == rank._plan_cost(plan).required - == 32229502659 + == 23331122883 ) selected = rank._select_next_micro_batch(requests, 0) assert ( @@ -154,7 +154,7 @@ def test_exact_pending_demand_survives_recovery(monkeypatch, pending_rank, fits_ plan = pending_rank._plan_flat_forward(requests) assert pending_rank._estimate_flat_forward(requests) is None assert g.plan_floor(pending_rank, plan) == (8296857600, 12705630112) - assert pending_rank._memory_check(plan).estimated_required_bytes == 32229502659 + assert pending_rank._memory_check(plan).estimated_required_bytes == 23331122883 _check_component_demand_recovery( monkeypatch, pending_rank, requests, fits_after=fits_after ) @@ -172,8 +172,8 @@ def test_original_installed_norm_preserves_pending_floor(layer): assert g.model_shapes(rank) is not None plan = rank._plan_flat_forward(full_requests()) assert g.plan_floor(rank, plan) == (8296857600, 12705630112) - assert rank._memory_check(plan).estimated_required_bytes == 32229502659 - assert rank._plan_cost(plan).required == 32229502659 + assert rank._memory_check(plan).estimated_required_bytes == 23331122883 + assert rank._plan_cost(plan).required == 23331122883 assert rank._estimate_flat_forward(full_requests()) is None for requests in ([], full_requests(no_grad=True)): assert g.plan_floor(rank, rank._plan_flat_forward(requests)) == (0, 0) diff --git a/tests/unit/test_trainer_rank_shared_memory.py b/tests/unit/test_trainer_rank_shared_memory.py index 6867e3566..7d08c8aa5 100644 --- a/tests/unit/test_trainer_rank_shared_memory.py +++ b/tests/unit/test_trainer_rank_shared_memory.py @@ -1,5 +1,6 @@ """One supported shared return held across routed compute; not all backward saves.""" +import math from types import SimpleNamespace import pytest @@ -126,7 +127,8 @@ def test_shared_return_in_actual_constructor_and_plan(layer, gate, no_grad): 8296857600, 50640 * (checkpoint_coefficient + 128) + 3157761952, ) - assert rank._plan_cost(plan).required == (32685829827 if gate else 32457666243) + # The incoming gradient replaces one gradient per boundary. + assert rank._plan_cost(plan).required == (23787450051 if gate else 23559286467) selected = rank._select_next_micro_batch(requests, 0) assert ( selected.check.estimated_required_bytes @@ -147,7 +149,7 @@ def test_original_norm_installation_preserves_shared_return(layer, gated): 8296857600, 50640 * (checkpoint_coefficient + 128) + 3157761952, ) - expected = 32685829827 if gated else 32457666243 + expected = 23787450051 if gated else 23559286467 assert rank._memory_check(plan).estimated_required_bytes == expected assert rank._plan_cost(plan).required == expected @@ -305,14 +307,18 @@ def test_pre_gate_cache_precedes_owned_dispatcher_and_is_checkpoint_only(layer): ) # Installed dispatcher partials must not be repriced. assert rank._moe_checkpoint_grad_bytes_per_token == 196608 groups = ((19, True), (23, False)) + # Gradient rows also keep the recomputed layer's residual and norm rows and + # its MoE routing state; the first call allocates TE's cuBLAS workspaces. + beside = 2 * 2048 * 2 + rank._moe_checkpoint_state_bytes_per_token() assert rank._checkpoint_memory_floor(groups) == ( 19 * 40 * 4096, max( - 19 * (196608 + 128), + 19 * (196608 + 128 + beside), 23 * (192512 + 4 * 2048 * 2), - 19 * (196608 - 32768 + 128) + 10485760, + 19 * (196608 - 32768 + 128 + beside) + 10485760, 23 * (192512 - 32768 + 128 + 4 * 2048 * 2) + 10485760, - ), + ) + + rank._te_workspace_growth_bytes(), ) for mode in (None, "selective"): rank.runtime.model[0].decoder.config.recompute_granularity = mode @@ -340,7 +346,7 @@ def test_pre_gate_mixed_reference_and_exact_cost_mode_selection(layer, gradient_ ) assert rank._checkpoint_memory_floor(rank._plan_group_rows(mixed)) == ( 67 * 40 * 4096, - 4096 * (192512 + 4 * 2048 * 2), + 4096 * (192512 + 4 * 2048 * 2) + rank._te_workspace_growth_bytes(), ) # A reference-only path must not read or validate the unused gradient cache. rank._moe_checkpoint_grad_bytes_per_token = None @@ -410,12 +416,12 @@ def test_shared_return_escapes_the_ep_routed_allowance(layer, checkpoint_grad, s layer.config.context_parallel_size = layer.config.expert_model_parallel_size = 2 # EP2 halves the experts each rank owns. _hybridep(layer, 2).token_dispatcher.num_local_experts = local_experts // 2 - # HybridEP's 1.5x allowance turns top-k 8 into 12 routed rows, each with - # one dispatched H-wide input; the gated shared return (doubled for + # HybridEP's EP2 allowance (1.4) turns top-k 8 into 11.2 routed rows, each + # with one dispatched H-wide input; the gated shared return (doubled for # checkpoint backward) is per local token. assert ( _moe_output_bytes_per_token( [layer], ParallelShape(tp=1, cp=2, ep=2), checkpoint_grad=checkpoint_grad ) - == 12 * 9728 * 2 + shared + == math.ceil(8 * 1.4 * 9728 * 2) + shared ) From 9d03f33481e5ac21c723ab426613a90154332641 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 20:02:09 +0000 Subject: [PATCH 09/40] Price HybridEP routed rows on the EP group's balanced share HybridEP dispatches the whole EP group's rows. When that group is this rank's CP group, a balanced rank receives the group's rows over EP, not the busiest CP rank's share: a CP2/EP2 real-data trace put 52,480 rows on one rank while each layer dispatched exactly 8 x 96,794 pairs across both. Price only the routed part (and its converted stages) on that share; the shared expert, mixer and boundaries stay on local rows. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/megatron/context_parallel/runtime.py | 33 ++++ src/art/trainer_rank/_impl.py | 167 ++++++++++++++++-- .../test_trainer_rank_admission_inputs.py | 1 + .../test_trainer_rank_checkpoint_memory.py | 26 +++ tests/unit/test_trainer_rank_moe_memory.py | 87 ++++++++- tests/unit/test_trainer_rank_shared_memory.py | 10 +- 6 files changed, 308 insertions(+), 16 deletions(-) diff --git a/src/art/megatron/context_parallel/runtime.py b/src/art/megatron/context_parallel/runtime.py index 821b60e76..9d1c8a216 100644 --- a/src/art/megatron/context_parallel/runtime.py +++ b/src/art/megatron/context_parallel/runtime.py @@ -433,6 +433,39 @@ def context_parallel_rank_model_token_counts( ) +def context_parallel_model_token_total( + *, + group_ids: torch.Tensor, + parent_ids: torch.Tensor, + topology: ParallelTopology, + config: ContextParallelConfig, + original_seq_len: int, + build_gdn_execution_spec: bool, + gdn_planner_config: Any | None = None, +) -> int: + """Return the CP group's model rows in its larger physical layout.""" + planning_key, bundle, _group_ids_cpu, _parent_ids_cpu = ( + _get_or_build_planning_bundle( + group_ids=group_ids, + parent_ids=parent_ids, + topology=topology, + config=config, + original_seq_len=original_seq_len, + build_gdn_execution_spec=build_gdn_execution_spec, + ) + ) + total = sum(bundle.token_layout_index.token_counts_by_rank) + if not build_gdn_execution_spec: + return total + decision = _plan_gdn_global_execution( + planning_key=planning_key, + bundle=bundle, + topology=topology, + gdn_planner_config=gdn_planner_config, + ) + return max(total, sum(decision.gdn_token_counts_by_rank)) + + def _normalized_chunk_size( *, valid_tokens: int, diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 222b6a371..f5df6fb22 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1571,6 +1571,29 @@ def _moe_dispatcher_supported( ) +def _ep_group_is_cp_group(shape: ParallelShape) -> bool: + """Whether this rank's expert-parallel group is exactly its CP group. + + HybridEP then dispatches that CP group's rows across it: at balanced + routing each rank receives the group's rows over EP, however CP split them. + """ + if shape.ep <= 1 or shape.ep != shape.cp or (shape.tp, shape.etp) != (1, 1): + return False + if not dist.is_available() or not dist.is_initialized(): + return False + try: + from megatron.core import parallel_state as ps + except ModuleNotFoundError: + return False + expert = ps.get_expert_model_parallel_group(check_initialized=False) + context = ps.get_context_parallel_group(check_initialized=False) + if expert is None or context is None: + return False + return sorted(dist.get_process_group_ranks(expert)) == sorted( + dist.get_process_group_ranks(context) + ) + + def _hybridep_rows_per_rank(capacity: int, ranks: int) -> int: """HybridEP's allocated rows per rank: TMA-aligned, at least 512, padded to the 64-row combine chunk.""" @@ -1602,8 +1625,13 @@ def _moe_output_bytes_per_token( checkpoint_grad: bool = False, converted_stages: list[tuple[int, int]] | None = None, slot_ref: "LoRASlotRef | None" = None, + shared_bytes: list[int] | None = None, ) -> int: - """Known routed-expert working set, not a complete model/compiled bound.""" + """Known routed-expert working set, not a complete model/compiled bound. + + ``shared_bytes`` collects each layer's shared-expert part of the per-token + coefficient and stages; that part follows local rows, not routed rows. + """ # CP shards rows, not the per-token working set. At EP>1 only ART's # HybridEP flex dispatcher is modeled; TP and ETP are not. if (shape.tp, shape.etp) != (1, 1): @@ -1746,6 +1774,8 @@ def _moe_output_bytes_per_token( # Gate-score backward saves a distinct pre-gate X. Charge it # beside this layer's returned X, not another layer's maximum. shared += shared + if shared_bytes is not None: + shared_bytes.append(shared) routed_rows = config.moe_router_topk * routed_allowance row_bytes = ( math.ceil(routed_rows * features * weights.element_size()) + shared @@ -1977,9 +2007,14 @@ def memory_field(name: str, default: Any = None) -> Any: ) forward_stages: list[tuple[int, int]] = [] gradient_stages: list[tuple[int, int]] = [] + forward_shared: list[int] = [] + gradient_shared: list[int] = [] self._moe_output_bytes_per_token = ( _moe_output_bytes_per_token( - runtime.model, self._parallel_shape, converted_stages=forward_stages + runtime.model, + self._parallel_shape, + converted_stages=forward_stages, + shared_bytes=forward_shared, ) if self._moe_layers else 0 @@ -1993,6 +2028,7 @@ def memory_field(name: str, default: Any = None) -> Any: self._parallel_shape, checkpoint_grad=True, converted_stages=gradient_stages, + shared_bytes=gradient_shared, ) if self._moe_layers else 0 @@ -2004,6 +2040,15 @@ def memory_field(name: str, default: Any = None) -> Any: self._moe_gradient_stages = ( tuple(gradient_stages) if self._moe_checkpoint_grad_bytes_per_token else () ) + self._moe_forward_shared_bytes = ( + max(forward_shared, default=0) if self._moe_output_bytes_per_token else 0 + ) + self._moe_gradient_shared_bytes = ( + max(gradient_shared, default=0) + if self._moe_checkpoint_grad_bytes_per_token + else 0 + ) + self._ep_group_is_cp_group = _ep_group_is_cp_group(self._parallel_shape) selection = select_scoring( device_capability=capability, device_memory_bytes=device_memory, @@ -3982,6 +4027,34 @@ def _plan_group_rows(self, plan: _FlatForwardPlan) -> tuple[tuple[int, bool], .. for group in plan.groups ) + def _plan_group_routed_rows(self, plan: _FlatForwardPlan) -> tuple[int, ...]: + """Rows one rank's experts receive per group at balanced routing. + + HybridEP dispatches the whole EP group's rows. When that group is this + rank's CP group, a balanced rank receives the group's rows over EP, + however unevenly CP split them; otherwise keep the local rows. + """ + rows = self._plan_group_rows(plan) + if ( + not getattr(self, "_ep_group_is_cp_group", False) + or plan.signature.topology[2] <= 1 + ): + return tuple(local for local, _ in rows) + topology = self._topology() + return tuple( + min( + local, + -( + -self._cp_group_model_tokens( + _pad_packed_batch(group.packed, multiple=int(topology.tp)), + topology=topology, + ) + // int(topology.cp) + ), + ) + for (local, _), group in zip(rows, plan.groups, strict=True) + ) + def _checkpoint_moe_bytes_per_token(self) -> int: forward = self._moe_output_bytes_per_token gradient = self._moe_checkpoint_grad_bytes_per_token @@ -3998,6 +4071,7 @@ def _moe_workspace_bytes( self, rows: int, *, + routed_rows: int | None = None, checkpoint_grad: bool = False, slot_ref: "LoRASlotRef | None" = None, ) -> int: @@ -4006,7 +4080,9 @@ def _moe_workspace_bytes( The constructor cache covers original tensors. Explicit slots are repriced from their tensor metadata and original owners, including this rank's exact dispatcher wrapper. Ordinary non-checkpoint gradients - retain only forward-stage coverage. + retain only forward-stage coverage. ``routed_rows`` is what this rank's + experts receive at balanced routing (``rows`` by default); the shared + expert's part stays on the local rows. """ coefficient = ( self._checkpoint_moe_bytes_per_token() @@ -4018,8 +4094,16 @@ def _moe_workspace_bytes( "_moe_gradient_stages" if checkpoint_grad else "_moe_forward_stages", (), ) + shared = getattr( + self, + "_moe_gradient_shared_bytes" + if checkpoint_grad + else "_moe_forward_shared_bytes", + 0, + ) if slot_ref is not None and slot_ref.name is not None: selected: list[tuple[int, int]] = [] + slot_shared: list[int] = [] coefficient = ( _moe_output_bytes_per_token( self.runtime.model, @@ -4027,11 +4111,13 @@ def _moe_workspace_bytes( checkpoint_grad=checkpoint_grad, converted_stages=selected, slot_ref=slot_ref, + shared_bytes=slot_shared, ) if self._moe_layers else 0 ) stages = tuple(selected) if coefficient else () + shared = max(slot_shared, default=0) if coefficient else 0 if type(stages) is not tuple or any( type(stage) is not tuple or len(stage) != 2 @@ -4039,22 +4125,28 @@ def _moe_workspace_bytes( for stage in stages ): raise ValueError("Invalid constructor converted-weight stages") - return ( + if type(shared) is not int or not 0 <= shared <= coefficient: + raise ValueError("Invalid constructor shared-expert coefficient") + routed = rows if routed_rows is None else max(0, min(rows, routed_rows)) + return (rows - routed) * shared + ( max( - rows * coefficient, - *(rows * per_row + fixed for per_row, fixed in stages), + routed * coefficient, + *(routed * per_row + fixed for per_row, fixed in stages), ) if stages and rows > 0 - else rows * coefficient + else routed * coefficient ) def _checkpoint_memory_floor( self, group_rows: tuple[tuple[int, bool], ...], slot_refs: tuple["LoRASlotRef | None", ...] | None = None, + routed_rows: tuple[int, ...] | None = None, ) -> tuple[int, int]: """Conservative saved-boundary charge and one recomputed layer's workspace. + ``routed_rows`` are each group's balanced dispatched rows per rank + (``_plan_group_routed_rows``); by default, its local rows. 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. @@ -4134,10 +4226,15 @@ def _checkpoint_memory_floor( if gradient_rows else 0 ) + routed = (None,) * len(group_rows) if routed_rows is None else routed_rows workspace = max( - self._moe_workspace_bytes(rows, checkpoint_grad=grad, slot_ref=ref) + self._moe_workspace_bytes( + rows, routed_rows=dispatched, checkpoint_grad=grad, slot_ref=ref + ) + (mixer * rows if grad else 4 * rows * self._hidden_size * 2) - for (rows, grad), ref in zip(group_rows, refs, strict=True) + for (rows, grad), ref, dispatched in zip( + group_rows, refs, routed, strict=True + ) ) if moe: workspace += self._te_workspace_growth_bytes() @@ -4203,6 +4300,7 @@ def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: logical_tokens=plan.active_logical_tokens, gdn_segments=plan.grad_segment_count, group_rows=self._plan_group_rows(plan), + group_routed_rows=self._plan_group_routed_rows(plan), slot_refs=tuple(g.slot_ref for g in plan.groups), head_workspace_bytes=self._plan_head_workspace_bytes(plan), checkpoint_floor=_gdn_memory.plan_floor(self, plan), @@ -4219,6 +4317,7 @@ def _subforward_cost( logical_tokens: int, gdn_segments: int = 0, group_rows: tuple[tuple[int, bool], ...] = (), + group_routed_rows: tuple[int, ...] | None = None, slot_refs: tuple["LoRASlotRef | None", ...] | None = None, head_workspace_bytes: int = 0, checkpoint_floor: tuple[int, int] = (0, 0), @@ -4232,6 +4331,7 @@ def _subforward_cost( logical_tokens=logical_tokens, gdn_segments=gdn_segments, group_rows=group_rows, + group_routed_rows=group_routed_rows, slot_refs=slot_refs, head_workspace_bytes=head_workspace_bytes, checkpoint_floor=checkpoint_floor, @@ -4239,7 +4339,7 @@ def _subforward_cost( include_checkpoint_input_gradient=False, ) checkpoint_retained, checkpoint_workspace = self._checkpoint_memory_floor( - group_rows, slot_refs + group_rows, slot_refs, group_routed_rows ) retained = self._retained_memory_bytes( signature, @@ -6335,6 +6435,7 @@ def _fill_planner_snapshot( "gdn_segments": child.grad_segment_count, "retained_tokens": self._plan_retained_tokens(child), "group_rows": self._plan_group_rows(child), + "group_routed_rows": self._plan_group_routed_rows(child), "hybridep_growth_bytes": ( self._plan_hybridep_growth_bytes(child) ), @@ -7059,6 +7160,7 @@ def _memory_check( logical_tokens=forward.active_logical_tokens, gdn_segments=forward.grad_segment_count, group_rows=self._plan_group_rows(forward), + group_routed_rows=self._plan_group_routed_rows(forward), slot_refs=tuple(g.slot_ref for g in forward.groups), head_workspace_bytes=self._plan_head_workspace_bytes(forward), checkpoint_floor=_gdn_memory.plan_floor(self, forward), @@ -7905,6 +8007,7 @@ def _estimate_required_memory_bytes_from_values( logical_tokens: int | None = None, gdn_segments: int = 0, group_rows: tuple[tuple[int, bool], ...] = (), + group_routed_rows: tuple[int, ...] | None = None, slot_refs: tuple["LoRASlotRef | None", ...] | None = None, head_workspace_bytes: int = 0, checkpoint_floor: tuple[int, int] = (0, 0), @@ -7982,19 +8085,23 @@ def _estimate_required_memory_bytes_from_values( ) # Groups execute sequentially: summed packed rows conservatively bound # this FC2 component, not all workspace or retained graphs. + grouped = signature.topology[2] > 1 and bool(group_rows) static_compute = max( static_compute, *( self._moe_workspace_bytes( - sum(rows for rows, _ in group_rows) - if signature.topology[2] > 1 and group_rows - else packed_tokens, + sum(rows for rows, _ in group_rows) if grouped else packed_tokens, + routed_rows=sum(group_routed_rows) + if grouped and group_routed_rows is not None + else None, slot_ref=ref, ) for ref in (slot_refs or (None,)) ), ) - retained, workspace = self._checkpoint_memory_floor(group_rows, slot_refs) + retained, workspace = self._checkpoint_memory_floor( + group_rows, slot_refs, group_routed_rows + ) static_compute = max( static_compute, max(retained, checkpoint_floor[0]) @@ -8969,6 +9076,38 @@ def _max_rank_model_tokens( ) raise AssertionError("unreachable") + def _cp_group_model_tokens( + self, + batch: PrefixTreePack, + *, + topology: "ParallelTopology", + ) -> int: + """The CP group's model rows in its larger physical layout.""" + from art.megatron.context_parallel.runtime import ( + context_parallel_model_token_total, + ) + from art.megatron.training.microbatches import ( + _context_parallel_config_for_provider, + _gdn_planner_config_for_provider, + ) + + handler = self.runtime.model_support_handler + return context_parallel_model_token_total( + group_ids=batch.group_ids, + parent_ids=batch.parent_ids, + topology=topology, + config=_context_parallel_config_for_provider( + self.runtime.provider, + self.device, + handler, + ), + original_seq_len=int(batch.tokens.shape[1]), + build_gdn_execution_spec=handler.build_gdn_execution_spec, + gdn_planner_config=_gdn_planner_config_for_provider( + self.runtime.provider, handler + ), + ) + def _prepare_context_parallel_forward( self, batch: PrefixTreePack, diff --git a/tests/unit/test_trainer_rank_admission_inputs.py b/tests/unit/test_trainer_rank_admission_inputs.py index a79d198e1..3481bfd6d 100644 --- a/tests/unit/test_trainer_rank_admission_inputs.py +++ b/tests/unit/test_trainer_rank_admission_inputs.py @@ -32,6 +32,7 @@ def assert_plan_values(rank, plan, values): assert values["signature"] == plan.signature assert values["gdn_segments"] == plan.grad_segment_count assert values["group_rows"] == rank._plan_group_rows(plan) + assert values["group_routed_rows"] == rank._plan_group_routed_rows(plan) assert values["head_workspace_bytes"] == rank._plan_head_workspace_bytes(plan) assert values["checkpoint_floor"] == _gdn_memory.plan_floor(rank, plan) assert values["retained_tokens"] == rank._plan_retained_tokens(plan) diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index f73a46e73..22d4bd1b7 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -515,6 +515,32 @@ def test_no_grad_enclosure_config_guard(field, value): assert r._checkpoint_memory_floor(((11, False),)) == (0, 0) +def test_routed_rows_move_only_the_routed_moe_part(): + # A CP2/EP2 real-data trace: the busiest CP rank held 52,480 rows, while + # HybridEP dispatched 8 x 96,794 pairs per layer across both ranks, so a + # balanced rank receives 48,397 rows' pairs. Boundaries, the mixer and the + # shared expert stay on the local rows. + r = rank() + r._topology_key = lambda: (1, 1, 2, 1) + r._moe_gradient_shared_bytes = 8192 + local, routed = 52480, 48397 + retained, workspace = r._checkpoint_memory_floor(((local, True),)) + assert r._checkpoint_memory_floor(((local, True),), None, (local,)) == ( + retained, + workspace, + ) + fewer = r._checkpoint_memory_floor(((local, True),), None, (routed,)) + assert fewer[0] == retained + assert workspace - fewer[1] == (local - routed) * (188416 - 8192) + # Never more routed rows than local ones. + assert r._moe_workspace_bytes( + 10, routed_rows=20, checkpoint_grad=True + ) == r._moe_workspace_bytes(10, checkpoint_grad=True) + r._moe_gradient_shared_bytes = 188417 + with pytest.raises(ValueError, match="shared-expert"): + r._moe_workspace_bytes(10, checkpoint_grad=True) + + def qwen36_attention(r): # Qwen3.6-35B-A3B attention: 16 heads and 2 query groups of 256, gated. r._geometry = replace( diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 51d8373a5..4af7fc775 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -604,7 +604,9 @@ def held(rows): # Growth enters the forward and checkpoint peaks, not forward retention. rank._update_memory_profile(plan, 10**9, retained_bytes=10**8) monkeypatch.setattr( - rank, "_checkpoint_memory_floor", lambda rows, refs=None: (10**7, 10**6) + rank, + "_checkpoint_memory_floor", + lambda rows, refs=None, routed=None: (10**7, 10**6), ) grown = rank._plan_cost(plan) monkeypatch.setattr(rank, "_plan_hybridep_growth_bytes", lambda plan: 0) @@ -623,6 +625,89 @@ def held(rows): assert rank._plan_hybridep_growth_bytes(plan) == 0 +def test_converted_stages_follow_routed_rows_and_shared_stays_local(): + rank = _rank() + rank._moe_output_bytes_per_token = 1000 + rank._moe_forward_stages = ((1200, 50),) + rank._moe_forward_shared_bytes = 100 + assert rank._moe_workspace_bytes(10) == 10 * 1200 + 50 + assert rank._moe_workspace_bytes(10, routed_rows=6) == 6 * 1200 + 50 + 4 * 100 + assert rank._moe_workspace_bytes(10, routed_rows=0) == 50 + 10 * 100 + + +def test_ep_group_must_be_exactly_the_cp_group(monkeypatch): + ps = pytest.importorskip("megatron.core.parallel_state") + from art.trainer_rank import _impl + + shape = ParallelShape(tp=1, cp=2, ep=2) + for other in ( + replace(shape, ep=1), + replace(shape, ep=4), + replace(shape, tp=2), + replace(shape, etp=2), + ): + assert not _impl._ep_group_is_cp_group(other) + # Without initialized process groups there is nothing to compare. + assert not _impl._ep_group_is_cp_group(shape) + monkeypatch.setattr(_impl.dist, "is_initialized", lambda: True) + expert, context = object(), object() + ranks = {id(expert): [0, 1], id(context): [1, 0]} + monkeypatch.setattr(ps, "get_expert_model_parallel_group", lambda **_: expert) + monkeypatch.setattr(ps, "get_context_parallel_group", lambda **_: context) + monkeypatch.setattr(_impl.dist, "get_process_group_ranks", lambda g: ranks[id(g)]) + assert _impl._ep_group_is_cp_group(shape) + # EP spanning ranks outside this CP group sees other batches' rows. + ranks[id(context)] = [0, 2] + assert not _impl._ep_group_is_cp_group(shape) + monkeypatch.setattr(ps, "get_context_parallel_group", lambda **_: None) + assert not _impl._ep_group_is_cp_group(shape) + + +def test_routed_rows_use_the_cp_group_share_only_when_it_is_the_ep_group( + monkeypatch, +): + rank = _rank() + plan = rank._plan_flat_forward( + [ForwardInput(input_tokens=torch.arange(64), target_tokens=torch.arange(64))] + ) + assert rank._plan_group_routed_rows(plan) == (64,) + plan = replace(plan, signature=replace(plan.signature, topology=(1, 1, 2, 1))) + monkeypatch.setattr(rank, "_topology", lambda: SimpleNamespace(tp=1, cp=2)) + monkeypatch.setattr(rank, "_plan_group_rows", lambda plan: ((52480, True),)) + totals = [96794] + monkeypatch.setattr( + rank, "_cp_group_model_tokens", lambda batch, topology: totals[0] + ) + assert rank._plan_group_routed_rows(plan) == (52480,) + rank._ep_group_is_cp_group = True + assert rank._plan_group_routed_rows(plan) == (48397,) + totals[0] = 10**6 + assert rank._plan_group_routed_rows(plan) == (52480,) + + +def test_cp_group_total_is_its_larger_layout(monkeypatch): + runtime = pytest.importorskip("art.megatron.context_parallel.runtime") + bundle = SimpleNamespace( + token_layout_index=SimpleNamespace(token_counts_by_rank=(52480, 44314)) + ) + monkeypatch.setattr( + runtime, + "_get_or_build_planning_bundle", + lambda **_: ("key", bundle, None, None), + ) + monkeypatch.setattr( + runtime, + "_plan_gdn_global_execution", + lambda **_: SimpleNamespace(gdn_token_counts_by_rank=(48100, 48800)), + ) + values: dict[str, Any] = dict( + group_ids=None, parent_ids=None, topology=None, config=None, original_seq_len=0 + ) + total = runtime.context_parallel_model_token_total + assert total(**values, build_gdn_execution_spec=False) == 96794 + assert total(**values, build_gdn_execution_spec=True) == 96900 + + def test_split_charges_the_largest_hybridep_growth_beside_any_child_peak(): from art.trainer_rank._impl import _SubforwardCost diff --git a/tests/unit/test_trainer_rank_shared_memory.py b/tests/unit/test_trainer_rank_shared_memory.py index 7d08c8aa5..6b7841bf0 100644 --- a/tests/unit/test_trainer_rank_shared_memory.py +++ b/tests/unit/test_trainer_rank_shared_memory.py @@ -107,6 +107,9 @@ def test_shared_return_in_actual_constructor_and_plan(layer, gate, no_grad): assert rank._moe_output_bytes_per_token == 192512 checkpoint_coefficient = 196608 if gate else 192512 assert rank._moe_checkpoint_grad_bytes_per_token == checkpoint_coefficient + # The shared return is the part that stays on local rows under HybridEP. + assert rank._moe_forward_shared_bytes == 4096 + assert rank._moe_gradient_shared_bytes == (8192 if gate else 4096) shapes = g.model_shapes(rank) assert shapes is not None and shapes[1][0].moe_bytes_per_row == 192512 requests = full_requests(no_grad) @@ -419,9 +422,14 @@ def test_shared_return_escapes_the_ep_routed_allowance(layer, checkpoint_grad, s # HybridEP's EP2 allowance (1.4) turns top-k 8 into 11.2 routed rows, each # with one dispatched H-wide input; the gated shared return (doubled for # checkpoint backward) is per local token. + collected: list[int] = [] assert ( _moe_output_bytes_per_token( - [layer], ParallelShape(tp=1, cp=2, ep=2), checkpoint_grad=checkpoint_grad + [layer], + ParallelShape(tp=1, cp=2, ep=2), + checkpoint_grad=checkpoint_grad, + shared_bytes=collected, ) == math.ceil(8 * 1.4 * 9728 * 2) + shared ) + assert collected == [shared] From 8bc0c986c62713e3bb8bfb522bc76c9ba394184b Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 20:47:56 +0000 Subject: [PATCH 10/40] Keep the per-boundary gradient unless MoE covers all recompute Charge only the one incoming gradient when every decoder layer is a priced MoE layer that encloses its FC1 stage, for each gradient group's slot. A positive FC2-only coefficient, dense layers or a slot that reprices to zero keep one gradient per boundary, which also covers unpriced recompute work. Count an empty CP rank's padding row in the EP group's total: dispatch runs at least one row per rank, and that row is routed too. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/megatron/context_parallel/runtime.py | 12 +++- src/art/trainer_rank/_impl.py | 72 +++++++++++++++---- ...trainer_rank_checkpoint_gradient_memory.py | 53 +++++++++++++- .../test_trainer_rank_checkpoint_memory.py | 2 + tests/unit/test_trainer_rank_moe_memory.py | 3 + .../unit/test_trainer_rank_pending_memory.py | 6 +- 6 files changed, 127 insertions(+), 21 deletions(-) diff --git a/src/art/megatron/context_parallel/runtime.py b/src/art/megatron/context_parallel/runtime.py index 9d1c8a216..531cb8829 100644 --- a/src/art/megatron/context_parallel/runtime.py +++ b/src/art/megatron/context_parallel/runtime.py @@ -443,7 +443,11 @@ def context_parallel_model_token_total( build_gdn_execution_spec: bool, gdn_planner_config: Any | None = None, ) -> int: - """Return the CP group's model rows in its larger physical layout.""" + """Return the CP group's model rows in its larger physical layout. + + Dispatch runs at least one row on every rank; an empty rank's padding row + passes through the model too. + """ planning_key, bundle, _group_ids_cpu, _parent_ids_cpu = ( _get_or_build_planning_bundle( group_ids=group_ids, @@ -454,7 +458,9 @@ def context_parallel_model_token_total( build_gdn_execution_spec=build_gdn_execution_spec, ) ) - total = sum(bundle.token_layout_index.token_counts_by_rank) + total = sum( + max(1, count) for count in bundle.token_layout_index.token_counts_by_rank + ) if not build_gdn_execution_spec: return total decision = _plan_gdn_global_execution( @@ -463,7 +469,7 @@ def context_parallel_model_token_total( topology=topology, gdn_planner_config=gdn_planner_config, ) - return max(total, sum(decision.gdn_token_counts_by_rank)) + return max(total, sum(max(1, count) for count in decision.gdn_token_counts_by_rank)) def _normalized_chunk_size( diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index f5df6fb22..61aec462a 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1532,14 +1532,15 @@ def _expert_lora_weight_storage( return (transposes if a.shape[2] < 8 else 0, transposes, effective) -# Routed rows on the most loaded rank at EP>1, relative to balanced routing. +# Routed rows on the most loaded rank at EP>1, relative to its balanced share. # Expert-shard load is uneven per layer, from the router's expert preferences, -# and larger batches do not average it away. Pretrained Qwen3.6-35B-A3B on 3.5M -# tokens of retail agent trajectories: worst layer 1.21 at EP2, 1.41 at EP4, -# 1.63 at EP8 (1.88 on a small rollout sample); one production EP2 run saw -# 1.35. These samples bound what was measured, not all routing. Unmeasured EP -# sizes use the next measured one; above EP8 the allowance grows with log2(EP) -# up to EP itself (every pair on one rank). +# and larger batches do not average it away. Qwen3.6-35B-A3B on 3.5M tokens of +# retail agent trajectories, worst layer at EP2 / EP4 / EP8: pretrained 1.21 / +# 1.41 / 1.63, a trained policy 1.24 / 1.41 / 1.61; a small rollout sample +# reached 1.95 at EP8. One production EP2 run was inferred at 1.35. These +# samples bound what was measured, not all routing. Unmeasured EP sizes use the +# next measured one; above EP8 the allowance grows with log2(EP) up to EP +# itself (every pair on one rank). _EP_ROUTED_ROW_ALLOWANCE = {2: 1.4, 4: 1.6, 8: 2.0} @@ -1626,11 +1627,13 @@ def _moe_output_bytes_per_token( converted_stages: list[tuple[int, int]] | None = None, slot_ref: "LoRASlotRef | None" = None, shared_bytes: list[int] | None = None, + enclosed: list[bool] | None = None, ) -> int: """Known routed-expert working set, not a complete model/compiled bound. ``shared_bytes`` collects each layer's shared-expert part of the per-token coefficient and stages; that part follows local rows, not routed rows. + ``enclosed`` records, per MoE layer, whether its FC1 stage is priced too. """ # CP shards rows, not the per-token working set. At EP>1 only ART's # HybridEP flex dispatcher is modeled; TP and ETP are not. @@ -1764,6 +1767,8 @@ def _moe_output_bytes_per_token( # path. This is one stage, not a backward bound. features += dispatched * fc2.out_features + fc1.out_features enclosing_fc1 = fc1 + if enclosed is not None: + enclosed.append(enclosing_fc1 is not None) shared = _shared_expert_output_bytes_per_token(layer) if ( checkpoint_grad @@ -2009,6 +2014,7 @@ def memory_field(name: str, default: Any = None) -> Any: gradient_stages: list[tuple[int, int]] = [] forward_shared: list[int] = [] gradient_shared: list[int] = [] + gradient_enclosed: list[bool] = [] self._moe_output_bytes_per_token = ( _moe_output_bytes_per_token( runtime.model, @@ -2029,6 +2035,7 @@ def memory_field(name: str, default: Any = None) -> Any: checkpoint_grad=True, converted_stages=gradient_stages, shared_bytes=gradient_shared, + enclosed=gradient_enclosed, ) if self._moe_layers else 0 @@ -2048,6 +2055,17 @@ def memory_field(name: str, default: Any = None) -> Any: if self._moe_checkpoint_grad_bytes_per_token else 0 ) + # Recompute is covered only if every decoder layer is a priced MoE + # layer whose FC1 stage is enclosed; otherwise dense MLP or FC1 work + # the floor does not price keeps the per-boundary gradient allowance. + self._moe_gradient_enclosed = ( + tuple(gradient_enclosed) + if self._moe_checkpoint_grad_bytes_per_token + else () + ) + self._moe_recompute_covered = len( + self._moe_gradient_enclosed + ) == self._num_layers and all(self._moe_gradient_enclosed) self._ep_group_is_cp_group = _ep_group_is_cp_group(self._parallel_shape) selection = select_scoring( device_capability=capability, @@ -4273,22 +4291,46 @@ def _moe_checkpoint_state_bytes_per_token(self) -> int: ) return routing + 3 * geometry.moe_shared_expert_ffn * self._param_dtype_size + def _moe_recompute_covered_for(self, slot_ref: "LoRASlotRef | None") -> bool: + """Whether an explicit slot keeps the constructor's full MoE coverage.""" + if not getattr(self, "_moe_recompute_covered", False): + return False + if slot_ref is None or slot_ref.name is None: + return True + enclosed: list[bool] = [] + coefficient = _moe_output_bytes_per_token( + self.runtime.model, + self._parallel_shape, + checkpoint_grad=True, + slot_ref=slot_ref, + enclosed=enclosed, + ) + return coefficient > 0 and len(enclosed) == self._num_layers and all(enclosed) + def _checkpoint_input_gradient_bytes( - self, group_rows: tuple[tuple[int, bool], ...] + self, + group_rows: tuple[tuple[int, bool], ...], + slot_refs: tuple["LoRASlotRef | None", ...] | None = None, ) -> int: """Gradient rows live at the recomputed layer's peak. Backward recomputes the last layer first, so its peak meets every saved - boundary but only the one incoming gradient. Where the MoE stage is - priced (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. + boundary but only the one incoming gradient. Where the MoE stage covers + every layer's recompute, FC1 included (Qwen3.6-35B-A3B traces at CP1, + CP2/EP1 and EP2/CP2), charge that gradient. Elsewhere keep one gradient + per boundary: that allowance also covers dense MLP and other recompute + work the floor does not price. """ retained, _ = self._checkpoint_memory_floor(group_rows) if not retained: return 0 gradient_rows = sum(rows for rows, grad in group_rows if grad) - if self._checkpoint_moe_bytes_per_token(): + 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 + ): return gradient_rows * self._hidden_size * 2 return retained @@ -4354,7 +4396,7 @@ def _subforward_cost( ) # Input gradients live at the recomputed layer's peak; kept out of # forward retention, including the cold fallback above. - gradient = self._checkpoint_input_gradient_bytes(group_rows) + gradient = self._checkpoint_input_gradient_bytes(group_rows, slot_refs) checkpoint_retained = output_bytes + max( checkpoint_retained, checkpoint_floor[0] ) @@ -8107,7 +8149,7 @@ def _estimate_required_memory_bytes_from_values( max(retained, checkpoint_floor[0]) + max(workspace, head_workspace_bytes, checkpoint_floor[1]) + ( - self._checkpoint_input_gradient_bytes(group_rows) + self._checkpoint_input_gradient_bytes(group_rows, slot_refs) if include_checkpoint_input_gradient else 0 ), diff --git a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py index b4008f6c3..cc61cb739 100644 --- a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py @@ -9,8 +9,12 @@ import pytest from test_trainer_rank_checkpoint_memory import price, rank, requests -from test_trainer_rank_moe_memory import layer # noqa: F401 -from test_trainer_rank_pending_memory import full_requests, pending_rank # noqa: F401 +from test_trainer_rank_moe_memory import _enclosing_moe, layer # noqa: F401 +from test_trainer_rank_pending_memory import ( # noqa: F401 + full_requests, + pending_rank, + rank_with_moe, +) import torch from art.trainer_rank import ForwardInput @@ -45,6 +49,51 @@ def test_pending_cold_peak_does_not_become_forward_retention(pending_rank): assert warm.required == cost.required +BOUNDARY_GRADIENTS = 8 * 6330 * 40 * 2048 * 2 + + +def test_single_gradient_needs_every_layer_priced(layer): + # One MoE layer among 40 dense stand-ins: the floor does not price the dense + # layers' recompute, so one gradient per boundary stays. + r = rank_with_moe(_enclosing_moe(layer), stand_in=False)[0] + assert r._moe_gradient_enclosed == (True,) and not r._moe_recompute_covered + plan = r._plan_flat_forward(full_requests()) + assert r._plan_cost(plan).checkpoint_input_gradient == BOUNDARY_GRADIENTS + + +def test_single_gradient_needs_the_fc1_stage(layer): + # Without permute fusion the FC1 stage is not enclosed: a positive FC2-only + # coefficient does not cover recompute. + moe = _enclosing_moe(layer) + moe.config.moe_permute_fusion = False + r = rank_with_moe(moe)[0] + assert r._checkpoint_moe_bytes_per_token() > 0 + assert r._moe_gradient_enclosed == (False,) and not r._moe_recompute_covered + plan = r._plan_flat_forward(full_requests()) + assert r._plan_cost(plan).checkpoint_input_gradient == BOUNDARY_GRADIENTS + + +def test_a_slot_that_loses_moe_coverage_keeps_boundary_gradients( + pending_rank, monkeypatch +): + from art.megatron.lora import LoRASlotRef + from art.trainer_rank import _impl + + r = pending_rank + ref = LoRASlotRef(kind="checkpoint", name="policy") + groups = ((100, True),) + assert r._checkpoint_input_gradient_bytes(groups) == 100 * 2048 * 2 + monkeypatch.setattr(_impl, "_moe_output_bytes_per_token", lambda *a, **k: 0) + assert r._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 40 * 4096 + + def covered(*args, enclosed, **kwargs): + enclosed.extend([True] * r._num_layers) + return 1 + + monkeypatch.setattr(_impl, "_moe_output_bytes_per_token", covered) + assert r._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 2048 * 2 + + @pytest.mark.parametrize("moe", [True, False]) @pytest.mark.parametrize("rows", [1, 67, 1024]) def test_attention_only_extent_scales_with_gradient_rows(rows, moe): diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index 22d4bd1b7..1d1087ab6 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -58,6 +58,8 @@ def rank(): ) result._moe_output_bytes_per_token = 188416 result._moe_checkpoint_grad_bytes_per_token = 188416 + # Stands in for Qwen3.6-35B-A3B, whose MoE stage prices every layer's recompute. + result._moe_recompute_covered = True return result diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 4af7fc775..d4e390ccf 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -706,6 +706,9 @@ def test_cp_group_total_is_its_larger_layout(monkeypatch): total = runtime.context_parallel_model_token_total assert total(**values, build_gdn_execution_spec=False) == 96794 assert total(**values, build_gdn_execution_spec=True) == 96900 + # An empty rank still dispatches, and routes, one padding row. + bundle.token_layout_index.token_counts_by_rank = (2, 0) + assert total(**values, build_gdn_execution_spec=False) == 3 def test_split_charges_the_largest_hybridep_growth_beside_any_child_peak(): diff --git a/tests/unit/test_trainer_rank_pending_memory.py b/tests/unit/test_trainer_rank_pending_memory.py index 0d53d1f10..98443af42 100644 --- a/tests/unit/test_trainer_rank_pending_memory.py +++ b/tests/unit/test_trainer_rank_pending_memory.py @@ -21,7 +21,7 @@ def module(cls): return obj -def rank_with_moe(moe_layer, *, install_hooks=False): +def rank_with_moe(moe_layer, *, install_hooks=False, stand_in=True): from megatron.core.ssm.gated_delta_net import GatedDeltaNet from megatron.core.transformer.transformer_block import TransformerBlock from transformer_engine.pytorch import RMSNorm @@ -96,6 +96,10 @@ def rank_with_moe(moe_layer, *, install_hooks=False): ) ) r._dp_rank_and_size = lambda: (0, 1) # Uninitialized MCore has no CPU DP group. + if stand_in: + # The one MoE layer stands in for all 40 of Qwen3.6-35B-A3B's: recompute + # is covered if that layer prices its FC1 stage too. + r._moe_recompute_covered = r._moe_gradient_enclosed == (True,) return r, gd From 01f89b24ff56197692fe40796045d9807369c579 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 21:07:54 +0000 Subject: [PATCH 11/40] Require priced FC1 stages before trusting a slot's MoE coverage A layer counts as enclosed only if its FC1 converted stages are priced too, unless FC1 has no adapter or the selected slot has no FC1 tensors. A slot with FC1 adapters but no FC2 adapter prices FC2 rows from the original metadata yet skips the whole converted-stage block, so it now keeps one gradient per boundary. A slot's walk must enclose as many layers as the constructor's, which already matched every decoder layer. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 30 +++++++++++++++---- ...trainer_rank_checkpoint_gradient_memory.py | 2 +- tests/unit/test_trainer_rank_slot_memory.py | 16 ++++++++++ 3 files changed, 41 insertions(+), 7 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 61aec462a..06fc5325b 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1535,9 +1535,9 @@ def _expert_lora_weight_storage( # Routed rows on the most loaded rank at EP>1, relative to its balanced share. # Expert-shard load is uneven per layer, from the router's expert preferences, # and larger batches do not average it away. Qwen3.6-35B-A3B on 3.5M tokens of -# retail agent trajectories, worst layer at EP2 / EP4 / EP8: pretrained 1.21 / -# 1.41 / 1.63, a trained policy 1.24 / 1.41 / 1.61; a small rollout sample -# reached 1.95 at EP8. One production EP2 run was inferred at 1.35. These +# retail agent trajectories, worst layer in 200k-token batches at EP2 / EP4 / +# EP8: pretrained up to 1.22 / 1.40 / 1.62, a trained policy up to 1.24 / 1.41 / +# 1.61; a small rollout sample reached 1.95 at EP8. One production EP2 run was inferred at 1.35. These # samples bound what was measured, not all routing. Unmeasured EP sizes use the # next measured one; above EP8 the allowance grows with log2(EP) up to EP # itself (every pair on one rank). @@ -1767,8 +1767,6 @@ def _moe_output_bytes_per_token( # path. This is one stage, not a backward bound. features += dispatched * fc2.out_features + fc1.out_features enclosing_fc1 = fc1 - if enclosed is not None: - enclosed.append(enclosing_fc1 is not None) shared = _shared_expert_output_bytes_per_token(layer) if ( checkpoint_grad @@ -1787,6 +1785,7 @@ def _moe_output_bytes_per_token( ) coefficient = max(coefficient, row_bytes) storage = _expert_lora_weight_storage(lora, slot_ref) + fc1_stages = False if converted_stages is not None and storage is not None: padded, transposes, effective = storage saved_fc1, rank_fc1 = 0, 0 @@ -1815,6 +1814,7 @@ def _moe_output_bytes_per_token( and first_tensors[1].shape[2] == enclosing_fc1.out_features ): first_padding, first_transposes, first_rank = first + fc1_stages = True # FC1 retains the routed H inputs and its base O1 # while producing adapter O1. Its sum is not live yet. converted_stages.append( @@ -1889,6 +1889,18 @@ def _moe_output_bytes_per_token( + 2 * (experts_count + 1) * 4, ) ) + if enclosed is not None: + # FC1 is covered when priced beside FC2 and its converted + # stages are priced too, unless it has no adapter or the + # selected slot has no FC1 tensors to convert. + adapter = getattr(enclosing_fc1, "lora", None) + inactive = adapter is None or ( + slot_ref is not None + and type(adapter) is LoRA + and "_slot" not in vars(adapter) + and _slot_lora_tensors(adapter, slot_ref) is None + ) + enclosed.append(enclosing_fc1 is not None and (fc1_stages or inactive)) if converted_stages is not None: # The EP allowance gives fractional routed rows; round each stage up. converted_stages[:] = [ @@ -4302,10 +4314,16 @@ def _moe_recompute_covered_for(self, slot_ref: "LoRASlotRef | None") -> bool: self.runtime.model, self._parallel_shape, checkpoint_grad=True, + converted_stages=[], slot_ref=slot_ref, enclosed=enclosed, ) - return coefficient > 0 and len(enclosed) == self._num_layers and all(enclosed) + # The constructor's walk already matched every decoder layer. + return ( + coefficient > 0 + and len(enclosed) == len(getattr(self, "_moe_gradient_enclosed", ())) + and all(enclosed) + ) def _checkpoint_input_gradient_bytes( self, diff --git a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py index cc61cb739..6681e2589 100644 --- a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py @@ -87,7 +87,7 @@ def test_a_slot_that_loses_moe_coverage_keeps_boundary_gradients( assert r._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 40 * 4096 def covered(*args, enclosed, **kwargs): - enclosed.extend([True] * r._num_layers) + enclosed.extend([True] * len(r._moe_gradient_enclosed)) return 1 monkeypatch.setattr(_impl, "_moe_output_bytes_per_token", covered) diff --git a/tests/unit/test_trainer_rank_slot_memory.py b/tests/unit/test_trainer_rank_slot_memory.py index f15a33fae..0fd2ee866 100644 --- a/tests/unit/test_trainer_rank_slot_memory.py +++ b/tests/unit/test_trainer_rank_slot_memory.py @@ -245,3 +245,19 @@ def guarded(name, *args, **kwargs): monkeypatch.setattr(builtins, "__import__", guarded) assert rank._slot_memory_shapes(ref) == () + + +def test_partial_slot_with_unpriced_fc1_stages_keeps_boundary_gradients(layer): + # A slot with FC1 adapters but no FC2 adapter still prices FC2 rows from the + # original metadata, but its FC1 converted weights go unpriced: keep one + # gradient per boundary for it. + rank, _ = rank_with_moe(weights(layer, 8)) + ref = load_slot(rank, "partial", 8) + groups = ((100, True),) + assert rank._moe_recompute_covered_for(ref) + assert rank._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 2048 * 2 + del layer.experts.linear_fc2.lora._slot_keys[ref] + assert rank._moe_workspace_bytes(1, checkpoint_grad=True, slot_ref=ref) > 0 + assert not rank._moe_recompute_covered_for(ref) + assert rank._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 40 * 4096 + assert rank._checkpoint_input_gradient_bytes(groups) == 100 * 2048 * 2 From ae10b7e03674723e59cd46ef989f139fbbcf4428 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 21:14:42 +0000 Subject: [PATCH 12/40] Rewrap the EP allowance comment Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 06fc5325b..8dfc41239 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1537,10 +1537,10 @@ def _expert_lora_weight_storage( # and larger batches do not average it away. Qwen3.6-35B-A3B on 3.5M tokens of # retail agent trajectories, worst layer in 200k-token batches at EP2 / EP4 / # EP8: pretrained up to 1.22 / 1.40 / 1.62, a trained policy up to 1.24 / 1.41 / -# 1.61; a small rollout sample reached 1.95 at EP8. One production EP2 run was inferred at 1.35. These -# samples bound what was measured, not all routing. Unmeasured EP sizes use the -# next measured one; above EP8 the allowance grows with log2(EP) up to EP -# itself (every pair on one rank). +# 1.61; a small rollout sample reached 1.95 at EP8. One production EP2 run was +# inferred at 1.35. These samples bound what was measured, not all routing. +# Unmeasured EP sizes use the next measured one; above EP8 the allowance grows +# with log2(EP) up to EP itself (every pair on one rank). _EP_ROUTED_ROW_ALLOWANCE = {2: 1.4, 4: 1.6, 8: 2.0} From 3d1d18e90b5c5b1d807fa95ed3c7fc52e660866d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 23:27:01 +0000 Subject: [PATCH 13/40] Keep boundary gradients above CP2; price TE beside the combine extent Review of the rebased stack: the MoE one-gradient allowance was traced at CP1 and CP2 only, while above CP2 a rank runs more remote attention stages than the mixer's CP2 allowance prices, so larger CP keeps one gradient per boundary. #971's combine-extent floor now includes the TE cuBLAS workspaces that are live beside it, instead of taking the maximum after they were added to the stage. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 26 +++++++++++++------ ...trainer_rank_checkpoint_gradient_memory.py | 11 ++++++++ tests/unit/test_trainer_rank_moe_memory.py | 3 ++- 3 files changed, 31 insertions(+), 9 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index e80ee4362..0378c8c24 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4353,7 +4353,12 @@ def _checkpoint_memory_floor( rows = max(rows for rows, _ in group_rows) if any(ref() is not None for ref in self._pending_hybridep_graphs): rows = max(rows, self._hybridep_rows_high_water) - workspace = max(workspace, -(-rows // 4) * 4 * self._hidden_size * 2) + # The combine output and the TE workspaces are live together. + workspace = max( + workspace, + -(-rows // 4) * 4 * self._hidden_size * 2 + + self._te_workspace_growth_bytes(), + ) return retained, workspace def _sequence_parallel_checkpoint_floor( @@ -4489,19 +4494,24 @@ def _checkpoint_input_gradient_bytes( Backward recomputes the last layer first, so its peak meets every saved boundary but only the one incoming gradient. Where the MoE stage covers every layer's recompute, FC1 included (Qwen3.6-35B-A3B traces at CP1, - CP2/EP1 and EP2/CP2), charge that gradient. Elsewhere keep one gradient - per boundary: that allowance also covers dense MLP and other recompute - work the floor does not price. + CP2/EP1 and EP2/CP2), charge that gradient. Elsewhere, including above + CP2 (more remote attention stages than the mixer's CP2 allowance), keep + one gradient per boundary: that allowance also covers dense MLP and + other recompute work the floor does not price. """ retained, _ = 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._topology_key()[2] <= 2 + and self._checkpoint_moe_bytes_per_token() + and all( + self._moe_recompute_covered_for(ref) + for (_, grad), ref in zip(group_rows, refs, strict=True) + if grad + ) ): return gradient_rows * self._hidden_size * 2 return retained diff --git a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py index 6681e2589..7eabcf57e 100644 --- a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py @@ -94,6 +94,17 @@ def covered(*args, enclosed, **kwargs): assert r._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 2048 * 2 +@pytest.mark.parametrize("cp", [1, 2, 4]) +def test_one_moe_gradient_only_up_to_the_traced_cp2(pending_rank, monkeypatch, cp): + """Above CP2 a rank runs more remote attention stages than the mixer's CP2 + allowance prices; the per-boundary gradient allowance must cover them.""" + r = pending_rank + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, cp, 1)) + groups = ((100, True),) + one, every = 100 * 2048 * 2, 100 * 40 * 4096 + assert r._checkpoint_input_gradient_bytes(groups) == (one if cp <= 2 else every) + + @pytest.mark.parametrize("moe", [True, False]) @pytest.mark.parametrize("rows", [1, 67, 1024]) def test_attention_only_extent_scales_with_gradient_rows(rows, moe): diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 0522e5161..57b1d1aca 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -841,7 +841,8 @@ def test_hybridep_recompute_prices_fresh_dense_output_without_buffer_growth( ) assert rank._plan_hybridep_growth_bytes(plan) == 0 retained, workspace = rank._checkpoint_memory_floor(groups) - assert workspace == 218752 * 2048 * 2 == 896008192 + # The combine output, with the TE workspaces live beside it. + assert workspace == 218752 * 2048 * 2 + rank._te_workspace_growth_bytes() cost = rank._subforward_cost(**values) assert cost.required == int((8 + 2 * retained + workspace) * 1.1) assert cost.checkpoint_workspace == workspace # Maximum, not stage + output. From f0670b9ec637b8143f55f3a4a2b4873ca92f2306 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 01:50:04 +0000 Subject: [PATCH 14/40] Price the adapter gradients a recompute backward holds at its peak A full-recompute backward allocates each layer's adapter gradients as it passes, and the step's optimizer frees them. Recomputing layer i still holds the saved boundaries of layers 0..i, so a short first wave peaks at layer 0 with nearly every layer's gradients live; the checkpoint floor priced only the last layer's end (all boundaries). On Qwen3.6-35B-A3B CP2, 2k and 4k token first waves were admitted 11.5% and 2.0% under their peaks. The floor now adds the largest excess of pending gradients over released boundaries across the real layers, while a slot's gradients are unallocated, and an unprofiled wave's 64 MiB of first-execution transients. Co-Authored-By: Claude Opus 5.5 (1M context) --- .github/workflows/prek.yml | 2 + src/art/trainer_rank/_impl.py | 170 +++++++++++- ...st_trainer_rank_adapter_gradient_memory.py | 258 ++++++++++++++++++ ...trainer_rank_checkpoint_gradient_memory.py | 24 +- tests/unit/test_trainer_rank_head_memory.py | 12 +- tests/unit/test_trainer_rank_moe_memory.py | 13 +- .../unit/test_trainer_rank_pending_memory.py | 11 +- .../unit/test_trainer_rank_planner_reports.py | 3 + tests/unit/test_trainer_rank_shared_memory.py | 4 +- tests/unit/test_trainer_rank_tp_floor.py | 14 +- 10 files changed, 482 insertions(+), 29 deletions(-) create mode 100644 tests/unit/test_trainer_rank_adapter_gradient_memory.py diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 18a49a302..67a05d36c 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -233,6 +233,7 @@ jobs: tests/unit/test_trainer_rank_profile_warm.py \ tests/unit/test_trainer_rank_tp_floor.py \ tests/unit/test_trainer_rank_checkpoint_gradient_memory.py \ + tests/unit/test_trainer_rank_adapter_gradient_memory.py \ tests/unit/test_trainer_rank_slot_memory.py \ tests/unit/test_trainer_rank_moe_memory.py \ tests/unit/test_trainer_rank_head_memory.py \ @@ -279,6 +280,7 @@ jobs: --ignore=tests/unit/test_trainer_rank_profile_warm.py \ --ignore=tests/unit/test_trainer_rank_tp_floor.py \ --ignore=tests/unit/test_trainer_rank_checkpoint_gradient_memory.py \ + --ignore=tests/unit/test_trainer_rank_adapter_gradient_memory.py \ --ignore=tests/unit/test_trainer_rank_slot_memory.py \ --ignore=tests/unit/test_trainer_rank_moe_memory.py \ --ignore=tests/unit/test_trainer_rank_head_memory.py \ diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index d4f7bde0f..01d2953c5 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -125,6 +125,10 @@ class TopK: _MEMORY_SAFETY_FACTOR = 1.10 _MEMORY_RESERVE_FRACTION = 0.03 _HEAD_CHUNK_TOKENS = 512 +# An unprofiled full-recompute gradient wave's first execution keeps two fixed +# 32 MiB transients live at its peak (Qwen3.6-35B-A3B CP2: the RoPE frequencies +# and a frozen linear's output, at 2k to 20k tokens); warm waves do not. +_COLD_RECOMPUTE_TRANSIENT_BYTES = 64 * 2**20 _PLANNER_REFINEMENT_BUDGET = 2_000 _LAYOUT_SELECTION_CACHE_LIMIT = 64 @@ -1061,6 +1065,13 @@ class _SubforwardCost: # HybridEP buffer growth before the safety factor. It is in required, not # retained, and persists across a split, which charges the largest once. hybridep_growth: int = 0 + # Adapter gradients the recompute backward holds beyond the boundaries it + # has released (``_checkpoint_adapter_gradient_bytes``), before the safety + # factor. Split children training the same slots share them, so a split + # charges the largest once; ``..._slots`` identifies those slots within + # this process (a hash, 0 when none), keeping the cost JSON-serializable. + checkpoint_adapter_gradient: int = 0 + checkpoint_adapter_gradient_slots: int = 0 @property def ephemeral(self) -> int: @@ -3408,12 +3419,32 @@ def _split_required_memory(costs: Sequence[_SubforwardCost]) -> int: if any(cost.checkpoint_input_gradient for cost in costs): # The caller owns all returned graphs. A calibrated forward-retained # discount cannot replace the sum of their input-gradient extents. + # Children training the same slots share their adapter gradients, + # allocated once by whichever child's backward reaches a layer + # first; the largest child's extra covers any order. Different + # slots have disjoint gradients, charged per child. + adapter = [ + cost.checkpoint_adapter_gradient + for cost in costs + if cost.checkpoint_adapter_gradient + ] + shared = ( + len( + { + cost.checkpoint_adapter_gradient_slots + for cost in costs + if cost.checkpoint_adapter_gradient + } + ) + <= 1 + ) checkpoint = ( sum( cost.checkpoint_retained + cost.checkpoint_input_gradient for cost in costs ) + max(cost.checkpoint_workspace for cost in costs) + + ((max(adapter) if shared else sum(adapter)) if adapter else 0) + growth ) required = max(required, int(checkpoint * _MEMORY_SAFETY_FACTOR)) @@ -4213,6 +4244,88 @@ def _checkpoint_memory_floor( workspace = max(workspace, -(-rows // 4) * 4 * self._hidden_size * 2) return retained, workspace + def _pending_adapter_gradient_bytes( + self, refs: Iterable["LoRASlotRef"] + ) -> tuple[int, ...]: + """Local adapter gradient bytes the next backward allocates, per decoder layer. + + One entry per decoder layer, then one for parameters outside the + decoder. Backward reaches a parameter's highest decoder layer first, so + a parameter shared across layers counts there once; those outside the + decoder count as live throughout. Only unallocated gradients count: + within a step, later waves find the rest in the availability baseline. + Empty when none is pending. + """ + refs = tuple(dict.fromkeys(refs)) + if not refs or len(self.runtime.model) != 1: + return () + try: + from art.megatron.lora import LoRA + except ModuleNotFoundError as error: + if error.name != "megatron": + raise + return () + chunk = self.runtime.model[0] + try: + layers = _language_model(chunk).decoder.layers + except (AttributeError, RuntimeError): + return () + layer_of: dict[int, int] = {} + params: dict[int, torch.nn.Parameter] = {} + + def slot_params(module: torch.nn.Module) -> Iterator[torch.nn.Parameter]: + for child in module.modules(): + # A LoRA without slot tables holds no slot parameters. + if isinstance(child, LoRA) and "_slot_keys" in vars(child): + for ref in refs: + yield from child.lora_slot_params(ref) + + for index, layer in enumerate(layers): + for param in slot_params(layer): + params[id(param)] = param + layer_of[id(param)] = max(layer_of.get(id(param), -1), index) + for param in slot_params(chunk): + if id(param) not in params: + params[id(param)] = param + layer_of[id(param)] = len(layers) + sizes = [0] * (len(layers) + 1) + for param_id, param in params.items(): + if ( + param.requires_grad + and param.grad is None + and getattr(param, "main_grad", None) is None + ): + sizes[layer_of[param_id]] += param.numel() * param.element_size() + return tuple(sizes) if any(sizes) else () + + def _checkpoint_adapter_gradient_bytes( + self, slots: Iterable["LoRASlotRef"], boundaries: Sequence[int] + ) -> int: + """The recompute backward's adapter-gradient peak beyond released boundaries. + + Backward recomputes the last layer first. While it recomputes layer i it + still holds the saved boundaries of layers 0..i and every adapter + gradient allocated so far: those of layers i..L-1 (a layer allocates its + own during its backward) and any outside the decoder. The floor already + prices all L boundaries at once, so the extra peak is + max(0, max over i of gradients(i..) - boundaries(i+1..)). It is taken + over the real per-layer sizes, not a uniform-layer line. A short + first wave peaks at layer 0 (Qwen3.6-35B-A3B CP2: 830-900 MB of expert + LoRA gradients live at its peak), a long one at the last layer. + ``boundaries`` gives each decoder layer's saved-boundary bytes; + ``slots`` are the gradient groups' adapter slots. + """ + pending = self._pending_adapter_gradient_bytes(slots) + if not pending or len(pending) != len(boundaries) + 1: + return 0 + extra = gradients = pending[-1] + released = 0 + for index in range(len(boundaries) - 1, -1, -1): + gradients += pending[index] + extra = max(extra, gradients - released) + released += boundaries[index] + return extra + def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( packed_tokens=plan.packed_tokens, @@ -4275,18 +4388,39 @@ def _subforward_cost( # backing stores, nor a bound for compiler saves or other backward work. # Keep it out of forward retention, including the cold fallback above. gradient = checkpoint_retained + gradient_slots = frozenset( + ref + for (_, grad), ref in zip( + group_rows, slot_refs or (None,) * len(group_rows), strict=True + ) + if grad and ref is not None + ) + adapter_gradient = ( + self._checkpoint_adapter_gradient_bytes( + gradient_slots, (gradient // self._num_layers,) * self._num_layers + ) + if gradient + else 0 + ) checkpoint_retained = output_bytes + max( checkpoint_retained, checkpoint_floor[0] ) checkpoint_workspace = max( checkpoint_workspace, head_workspace_bytes, checkpoint_floor[1] ) + if gradient and self._memory_profiles.get(signature) is None: + checkpoint_workspace += _COLD_RECOMPUTE_TRANSIENT_BYTES forward_required = required if gradient: required = max( required, int( - (checkpoint_retained + checkpoint_workspace + gradient) + ( + checkpoint_retained + + checkpoint_workspace + + gradient + + adapter_gradient + ) * _MEMORY_SAFETY_FACTOR ), ) @@ -4300,6 +4434,10 @@ def _subforward_cost( checkpoint_input_gradient=gradient, checkpoint_peak_increment=required - forward_required, hybridep_growth=hybridep_growth_bytes, + checkpoint_adapter_gradient=adapter_gradient, + checkpoint_adapter_gradient_slots=hash(gradient_slots) + if adapter_gradient + else 0, ) def _retained_memory_bytes( @@ -6060,6 +6198,17 @@ def _estimate_flat_forward( # This cheap return type has no slot metadata. Materialize the # exact plan instead of admitting with the constructor rank. return None + gradient_slots = [ + ref for (ref, grad), _ in groups if grad and ref is not None + ] + if ( + gradient_slots + and getattr(self, "_recompute_granularity", None) == "full" + and self._pending_adapter_gradient_bytes(gradient_slots) + ): + # The step's first backward allocates adapter gradients that + # only slot metadata can price; the exact plan carries it. + return None if ( any(mode for (_, mode), _ in groups) and _gdn_memory.model_shapes(self) is not None @@ -8071,11 +8220,28 @@ def _estimate_required_memory_bytes_from_values( retained, workspace = self._checkpoint_memory_floor( group_rows, slot_refs, gdn_segments ) + backward = 0 + if include_checkpoint_input_gradient and retained: + # The backward's other end and cold transients, as _subforward_cost. + backward = retained + self._checkpoint_adapter_gradient_bytes( + ( + ref + for (_, grad), ref in zip( + group_rows, + slot_refs or (None,) * len(group_rows), + strict=True, + ) + if grad and ref is not None + ), + (retained // self._num_layers,) * self._num_layers, + ) + if profiled is None and any(grad for _, grad in group_rows): + backward += _COLD_RECOMPUTE_TRANSIENT_BYTES static_compute = max( static_compute, max(retained, checkpoint_floor[0]) + max(workspace, head_workspace_bytes, checkpoint_floor[1]) - + (retained if include_checkpoint_input_gradient else 0), + + backward, ) if signature.topology[2] > 1: # Local head results coexist with full CP outputs during gathering. diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py new file mode 100644 index 000000000..dad274d7d --- /dev/null +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -0,0 +1,258 @@ +"""Adapter gradients at the recompute backward's peak: CPU admission math. + +Qwen3.6-35B-A3B CP2 allocator traces: a short first wave peaks at layer 0 with +830-900 MB of expert LoRA gradients live, a long one at the last layer. +""" + +from collections.abc import Sequence +import random + +import pytest +from test_trainer_rank_checkpoint_memory import rank, requests +import torch + +from art.megatron.lora import LoRA, LoRASlotRef +from art.trainer_rank import TrainerRank +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) +from art.trainer_rank._impl import _MemoryProfile, _SubforwardCost + +POLICY = LoRASlotRef("checkpoint", "policy") +OTHER = LoRASlotRef("checkpoint", "other") + + +def oracle(pending, boundaries): + """Live bytes beyond the floor while backward recomputes each layer.""" + layers = len(boundaries) + return max( + 0, + *( + sum(pending[index:layers]) + pending[layers] - sum(boundaries[index + 1 :]) + for index in range(layers) + ), + ) + + +def with_pending(monkeypatch, r, pending): + # Like the real method: no slots, nothing pending. + monkeypatch.setattr( + r, + "_pending_adapter_gradient_bytes", + lambda refs: tuple(pending) if tuple(refs) else (), + ) + + +@pytest.mark.parametrize( + "pending, boundaries", + [ + # Uniform gradients above the boundaries: layer 0 is the peak. + ([23] * 4 + [0], [4] * 4), + # Uniform gradients below the boundaries: the last layer is the peak. + ([3] * 4 + [0], [10] * 4), + # Attention and GDN layers differ; the peak is interior to neither end. + ([0, 100, 0, 100, 5], [60, 60, 60, 60]), + # Uneven boundaries, gradients outside the decoder live throughout. + ([7, 0, 50, 1, 9], [1, 80, 2, 3]), + ], +) +def test_extra_is_the_worst_backward_layer_over_real_sizes( + monkeypatch, pending, boundaries +): + r = rank() + with_pending(monkeypatch, r, pending) + assert r._checkpoint_adapter_gradient_bytes((POLICY,), boundaries) == oracle( + pending, boundaries + ) + + +def test_extra_matches_every_layer_for_random_sizes(monkeypatch): + r = rank() + generator = random.Random(0) + for _ in range(200): + layers = generator.randint(1, 12) + pending = [ + generator.choice((0, generator.randint(0, 50))) for _ in range(layers + 1) + ] + boundaries = [generator.randint(0, 40) for _ in range(layers)] + with_pending(monkeypatch, r, pending) + extra = r._checkpoint_adapter_gradient_bytes((POLICY,), boundaries) + assert extra == oracle(pending, boundaries) + + +def test_no_pending_gradients_add_nothing(monkeypatch): + r = rank() + with_pending(monkeypatch, r, ()) + assert r._checkpoint_adapter_gradient_bytes((POLICY,), [4] * 40) == 0 + + +def lora(**slots: list[torch.nn.Parameter]) -> LoRA: + module = LoRA.__new__(LoRA) + torch.nn.Module.__init__(module) + module._slot_keys = {} + by_name = { + LoRASlotRef("checkpoint", name): params for name, params in slots.items() + } + module.lora_slot_params = lambda ref: by_name.get(ref, []) # type: ignore[method-assign] + return module + + +def parameter(elements: int, *, dtype=torch.bfloat16) -> torch.nn.Parameter: + return torch.nn.Parameter(torch.zeros(elements, dtype=dtype)) + + +def adapter_rank( + layers: Sequence[torch.nn.Module], outside: torch.nn.Module +) -> TrainerRank: + r = rank() + block = r.runtime.model[0].decoder + block.layers = torch.nn.ModuleList(layers) + r.runtime.model[0].head = outside + return r + + +def test_pending_gradients_count_local_unallocated_slot_parameters(): + shared = parameter(10) + allocated = parameter(1000) + allocated.grad = torch.zeros_like(allocated) + master = parameter(1000) + setattr(master, "main_grad", torch.zeros(1000)) + frozen = parameter(1000) + frozen.requires_grad_(False) + layers = [ + lora(policy=[parameter(3), shared]), + lora(policy=[parameter(5)], other=[parameter(1000)]), + lora(policy=[allocated, master, frozen]), + lora(policy=[shared, parameter(4, dtype=torch.float32)]), + ] + r = adapter_rank(layers, lora(policy=[parameter(6)])) + pending = r._pending_adapter_gradient_bytes([POLICY]) + # BF16 bytes per layer; the shared parameter counts once, at its highest + # layer; allocated, main-grad and frozen parameters are not pending; the + # parameter outside the decoder is live throughout. + assert pending == (3 * 2, 5 * 2, 0, (10 + 0) * 2 + 4 * 4, 6 * 2) + assert r._pending_adapter_gradient_bytes([POLICY, OTHER]) == ( + 3 * 2, + 5 * 2 + 1000 * 2, + 0, + 10 * 2 + 4 * 4, + 6 * 2, + ) + assert r._pending_adapter_gradient_bytes([LoRASlotRef("checkpoint", "x")]) == () + assert r._pending_adapter_gradient_bytes([]) == () + + +def test_a_step_with_allocated_gradients_prices_no_extra(): + params = [parameter(100) for _ in range(4)] + r = adapter_rank([lora(policy=[p]) for p in params], torch.nn.Module()) + assert r._pending_adapter_gradient_bytes([POLICY]) == (200, 200, 200, 200, 0) + for p in params: + p.grad = torch.zeros_like(p) + # Later waves of the step find them in the availability baseline. + assert r._pending_adapter_gradient_bytes([POLICY]) == () + + +def test_slotless_lora_modules_hold_no_slot_parameters(): + module = LoRA.__new__(LoRA) + torch.nn.Module.__init__(module) + r = adapter_rank([module], torch.nn.Module()) + assert r._pending_adapter_gradient_bytes([POLICY]) == () + + +def priced(r, values, slot_refs): + n, out, signature, groups, head = values + return r._subforward_cost( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=slot_refs, + head_workspace_bytes=head, + ) + + +def test_cost_and_estimate_charge_the_extra_while_gradients_are_pending(monkeypatch): + r = rank() + values = r._estimate_flat_forward(requests(67, 4096)) + n, out, signature, groups, head = values + retained, workspace = r._checkpoint_memory_floor(groups) + boundary = retained // 40 + pending = [23 * 2**20] * 40 + [0] + with_pending(monkeypatch, r, pending) + extra = oracle(pending, [boundary] * 40) + assert extra > 0 + cost = priced(r, values, (POLICY, None)) + assert cost.checkpoint_adapter_gradient == extra + assert cost.checkpoint_adapter_gradient_slots == hash(frozenset({POLICY})) + assert cost.required == int( + (out + retained + workspace + COLD + retained + extra) * 1.1 + ) + estimate = r._estimate_required_memory_bytes_from_values( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=(POLICY, None), + head_workspace_bytes=head, + ) + assert estimate == cost.required + # Only gradient groups' slots count; no slot, no extra. + assert priced(r, values, (None, POLICY)).checkpoint_adapter_gradient == 0 + with_pending(monkeypatch, r, ()) + plain = priced(r, values, (POLICY, None)) + assert plain.checkpoint_adapter_gradient == 0 + assert plain.required == int((out + 2 * retained + workspace + COLD) * 1.1) + # A learned profile drops the first-execution transients, not the extra. + with_pending(monkeypatch, r, pending) + r._memory_profiles[signature] = _MemoryProfile(bytes_per_token=1, packed_tokens=n) + assert priced(r, values, (POLICY, None)).required == int( + (out + 2 * retained + workspace + extra) * 1.1 + ) + + +def child(extra: int, slots: int, workspace: int = 10) -> _SubforwardCost: + return _SubforwardCost( + required=int((1 + workspace + 1 + extra) * 1.1), + retained=0, + checkpoint_retained=1, + checkpoint_workspace=workspace, + checkpoint_input_gradient=1, + checkpoint_adapter_gradient=extra, + checkpoint_adapter_gradient_slots=slots, + ) + + +def test_split_charges_shared_gradients_once_and_distinct_slots_each(): + shared = [child(500, 7), child(300, 7)] + assert TrainerRank._split_required_memory(shared) == int((2 + 2 + 10 + 500) * 1.1) + distinct = [child(500, 7), child(300, 9)] + assert TrainerRank._split_required_memory(distinct) == int((2 + 2 + 10 + 800) * 1.1) + # A child with no pending gradients does not change the shared charge. + assert TrainerRank._split_required_memory([child(500, 7), child(0, 0)]) == int( + (2 + 2 + 10 + 500) * 1.1 + ) + + +def test_cheap_estimate_defers_while_a_gradient_slot_has_pending_gradients( + monkeypatch, +): + r = rank() + monkeypatch.setattr(r, "_ensure_checkpoint_slots_for", lambda *a, **k: None) + monkeypatch.setattr( + r, + "_resolve_slot_ref", + lambda request, checkpoint: POLICY if not request.no_grad else None, + ) + seen: list[tuple[LoRASlotRef, ...]] = [] + + def pending(refs): + seen.append(tuple(refs)) + return (1,) * 41 + + monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", pending) + assert r._estimate_flat_forward(requests(67, 4096)) is None + assert seen == [(POLICY,)] + monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", lambda refs: ()) + assert r._estimate_flat_forward(requests(67, 4096)) is not None diff --git a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py index d6d9ded56..83c1ad50a 100644 --- a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py @@ -9,6 +9,9 @@ import torch from art.trainer_rank import ForwardInput +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) from art.trainer_rank._impl import Unset, _MemoryProfile, _SplitForwardPlan @@ -29,12 +32,15 @@ def test_pending_cold_peak_does_not_become_forward_retention(pending_rank): assert cost.checkpoint_input_gradient == gradient # Exact previous cold estimate, including outputs and its one safety factor. assert cost.retained == 23102959299 - assert cost.required == int((plan.output_bytes + 2 * gradient + 12705630112) * 1.1) + # Unprofiled: the first execution's transients beside the workspace. + assert cost.required == int( + (plan.output_bytes + 2 * gradient + 12705630112 + COLD) * 1.1 + ) assert r._memory_check(plan).estimated_required_bytes == cost.required profile(r, plan) warm = r._plan_cost(plan) assert warm.retained == int((plan.output_bytes + gradient) * 1.1) - assert warm.required == cost.required + assert warm.required == int((plan.output_bytes + 2 * gradient + 12705630112) * 1.1) @pytest.mark.parametrize("rows", [1, 67, 1024]) @@ -60,8 +66,8 @@ def test_gradient_is_not_absorbed_by_larger_head_workspace(): head = 10**10 cost = price(r, (n, out, sig, groups, head)) gradient = 67 * 40 * 2048 * 2 - assert cost.checkpoint_workspace == head - assert cost.required == int((out + head + 2 * gradient) * 1.1) + assert cost.checkpoint_workspace == head + COLD + assert cost.required == int((out + head + COLD + 2 * gradient) * 1.1) assert cost.retained == int((out + head + gradient) * 1.1) r._memory_profiles[sig] = _MemoryProfile( bytes_per_token=10**9, @@ -199,16 +205,18 @@ def test_split_priority_subtracts_only_uncovered_gradient_peak(fully_masked): r = rank() plan = r._plan_flat_forward(requests(17, 19)) cold = r._plan_cost(plan) + # Once profiled, the static estimate has no first-execution transients. + profile(r, plan) + static = r._plan_cost(plan).required + assert static < cold.required # Place a real learned peak between the two static estimates, or above both. - measured = ( - cold.required + 10**7 if fully_masked else (cold.retained + cold.required) / 2 - ) + measured = static + 10**7 if fully_masked else (cold.retained + static) / 2 rate = (measured / 1.1 - plan.output_bytes) / plan.packed_tokens profile(r, plan, rate=rate) cost = r._plan_cost(plan) old_required = int((plan.output_bytes + int(plan.packed_tokens * rate)) * 1.1) assert cold.retained < old_required - assert cost.required == max(cold.required, old_required) + assert cost.required == max(static, old_required) assert ( cost.ephemeral - cost.checkpoint_peak_increment == old_required - cost.retained ) diff --git a/tests/unit/test_trainer_rank_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py index 21b7ad071..54d279f91 100644 --- a/tests/unit/test_trainer_rank_head_memory.py +++ b/tests/unit/test_trainer_rank_head_memory.py @@ -8,6 +8,9 @@ import torch from art.trainer_rank import ForwardInput +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) from art.trainer_rank._impl import ( _PACKED_PRICED_LOGICAL_ROW_BYTES, Unset, @@ -147,7 +150,10 @@ def test_outputs_retention_and_empirical_peak_are_counted_once(): head = 3 * 512 * 248320 * 2 cost = r._plan_cost(plan) assert cost.retained == int((plan.output_bytes + retained + head) * 1.1) - assert cost.required == int((plan.output_bytes + retained + gradient + head) * 1.1) + # Unprofiled: the first execution's transients beside the head workspace. + assert cost.required == int( + (plan.output_bytes + retained + gradient + head + COLD) * 1.1 + ) r._memory_profiles[plan.signature] = _MemoryProfile( bytes_per_token=2_000_000, packed_tokens=512, @@ -312,8 +318,8 @@ def test_target_backward_refuses_budget_below_logits_and_both_gradients(rows): retained, _ = r._checkpoint_memory_floor(r._plan_group_rows(plan)) gradient = rows * 40 * 2048 * 2 dense = min(rows, 512) * 248320 * 2 - before = int((plan.output_bytes + retained + gradient + 2 * dense) * 1.1) - expected = int((plan.output_bytes + retained + gradient + 3 * dense) * 1.1) + before = int((plan.output_bytes + retained + gradient + 2 * dense + COLD) * 1.1) + expected = int((plan.output_bytes + retained + gradient + 3 * dense + COLD) * 1.1) r._available_memory_bytes = lambda: (before + expected) // 2 check = r._memory_check(plan) assert check.estimated_required_bytes == expected diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 3fbdcee89..9cb0315ba 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -9,6 +9,9 @@ import torch from art.trainer_rank import ForwardInput, TrainerRank +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) from art.trainer_rank._impl import ( _PACKED_PRICED_LOGICAL_ROW_BYTES, _MemoryProfile, @@ -741,15 +744,19 @@ def test_hybridep_recompute_prices_fresh_dense_output_without_buffer_growth( retained, workspace = rank._checkpoint_memory_floor(groups) assert workspace == 218752 * 2048 * 2 == 896008192 cost = rank._subforward_cost(**values) - assert cost.required == int((8 + 2 * retained + workspace) * 1.1) - assert cost.checkpoint_workspace == workspace # Maximum, not stage + output. + # Unprofiled: the first execution's transients sit beside the extent. + assert cost.required == int((8 + 2 * retained + workspace + COLD) * 1.1) + assert cost.checkpoint_workspace == workspace + COLD # Not stage + output. rank._available_memory_bytes = lambda: 600000000 assert rank._memory_check_required(baseline.required).fits assert not rank._memory_check_required(cost.required).fits rank._memory_profiles[signature] = _MemoryProfile( bytes_per_token=1, packed_tokens=2 ) - assert rank._subforward_cost(**values).required == cost.required + # Profiled: the floor alone, without first-execution transients. + assert rank._subforward_cost(**values).required == int( + (8 + 2 * retained + workspace) * 1.1 + ) rank._memory_profiles[signature] = _MemoryProfile( bytes_per_token=10**9, packed_tokens=2 ) diff --git a/tests/unit/test_trainer_rank_pending_memory.py b/tests/unit/test_trainer_rank_pending_memory.py index 359de8daa..8cae0619a 100644 --- a/tests/unit/test_trainer_rank_pending_memory.py +++ b/tests/unit/test_trainer_rank_pending_memory.py @@ -12,6 +12,7 @@ from art.megatron.prefix_tree_packing import prefix_tree_pack from art.trainer_rank import ForwardInput, TrainerRank from art.trainer_rank import _gdn_memory as g +from art.trainer_rank._impl import _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD from art.trainer_rank._impl import Unset, _MemoryProfile @@ -133,7 +134,7 @@ def test_actual_constructor_cache_and_full_plan(pending_rank): assert ( rank._memory_check(plan).estimated_required_bytes == rank._plan_cost(plan).required - == 32229502659 + == 32303322409 ) selected = rank._select_next_micro_batch(requests, 0) assert ( @@ -196,7 +197,7 @@ def test_exact_pending_demand_survives_recovery(monkeypatch, pending_rank, fits_ plan = pending_rank._plan_flat_forward(requests) assert pending_rank._estimate_flat_forward(requests) is None assert g.plan_floor(pending_rank, plan) == (8296857600, 12705630112) - assert pending_rank._memory_check(plan).estimated_required_bytes == 32229502659 + assert pending_rank._memory_check(plan).estimated_required_bytes == 32303322409 _check_component_demand_recovery( monkeypatch, pending_rank, requests, fits_after=fits_after ) @@ -214,8 +215,8 @@ def test_original_installed_norm_preserves_pending_floor(layer): assert g.model_shapes(rank) is not None plan = rank._plan_flat_forward(full_requests()) assert g.plan_floor(rank, plan) == (8296857600, 12705630112) - assert rank._memory_check(plan).estimated_required_bytes == 32229502659 - assert rank._plan_cost(plan).required == 32229502659 + assert rank._memory_check(plan).estimated_required_bytes == 32303322409 + assert rank._plan_cost(plan).required == 32303322409 assert rank._estimate_flat_forward(full_requests()) is None for requests in ([], full_requests(no_grad=True)): assert g.plan_floor(rank, rank._plan_flat_forward(requests)) == (0, 0) @@ -384,7 +385,7 @@ def test_constructor_declined_moe_keeps_generic_admission(layer, unsupported): required = rank._plan_cost(plan).required # Generic checkpoint-input accounting still applies without a MoE component. gradient = 50640 * 40 * 2048 * 2 - assert required == int((plan.output_bytes + 2 * gradient) * 1.1) + assert required == int((plan.output_bytes + 2 * gradient + COLD) * 1.1) rank._available_memory_bytes = lambda: required - 1 assert not rank._memory_check(plan).fits rank._available_memory_bytes = lambda: required diff --git a/tests/unit/test_trainer_rank_planner_reports.py b/tests/unit/test_trainer_rank_planner_reports.py index 1adce5472..22bddb218 100644 --- a/tests/unit/test_trainer_rank_planner_reports.py +++ b/tests/unit/test_trainer_rank_planner_reports.py @@ -298,6 +298,8 @@ def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): "checkpoint_input_gradient": 0, "checkpoint_peak_increment": 0, "hybridep_growth": 0, + "checkpoint_adapter_gradient": 0, + "checkpoint_adapter_gradient_slots": 0, }, } ], @@ -379,6 +381,7 @@ def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): "checkpoint_workspace", "checkpoint_peak_increment", "hybridep_growth", + "checkpoint_adapter_gradient", ): altered = json.loads(path.read_bytes()) altered["replay"]["memory_replay"]["estimates"][0]["cost_components"][ diff --git a/tests/unit/test_trainer_rank_shared_memory.py b/tests/unit/test_trainer_rank_shared_memory.py index 36bc971d7..d8b1a2938 100644 --- a/tests/unit/test_trainer_rank_shared_memory.py +++ b/tests/unit/test_trainer_rank_shared_memory.py @@ -126,7 +126,7 @@ def test_shared_return_in_actual_constructor_and_plan(layer, gate, no_grad): 8296857600, 50640 * (checkpoint_coefficient + 128) + 3157761952, ) - assert rank._plan_cost(plan).required == (32685829827 if gate else 32457666243) + assert rank._plan_cost(plan).required == (32759649577 if gate else 32531485993) selected = rank._select_next_micro_batch(requests, 0) assert ( selected.check.estimated_required_bytes @@ -147,7 +147,7 @@ def test_original_norm_installation_preserves_shared_return(layer, gated): 8296857600, 50640 * (checkpoint_coefficient + 128) + 3157761952, ) - expected = 32685829827 if gated else 32457666243 + expected = 32759649577 if gated else 32531485993 assert rank._memory_check(plan).estimated_required_bytes == expected assert rank._plan_cost(plan).required == expected diff --git a/tests/unit/test_trainer_rank_tp_floor.py b/tests/unit/test_trainer_rank_tp_floor.py index 6d5f4d00d..c3f49e3b9 100644 --- a/tests/unit/test_trainer_rank_tp_floor.py +++ b/tests/unit/test_trainer_rank_tp_floor.py @@ -14,6 +14,7 @@ import torch from art.trainer_rank import TrainerRank +from art.trainer_rank._impl import _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD from art.trainer_rank._impl import _MemorySignature H, F, LAYERS = 5120, 17408, 64 @@ -114,10 +115,11 @@ def test_the_traced_tp4_wave_prices_its_boundary_shards_and_their_repeat(): assert workspace == 3 * SEGMENT cost = _required(r) assert cost.checkpoint_input_gradient == retained - # One segment plus up to three TP-padding roots, each with its states. - state = cost.checkpoint_workspace + # One segment plus up to three TP-padding roots, each with its states, + # beside the unprofiled first execution's transients. + state = cost.checkpoint_workspace - COLD assert state == 4 * SEGMENT - assert cost.required == int((OUTPUT + 2 * retained + state) * 1.1) + assert cost.required == int((OUTPUT + 2 * retained + state + COLD) * 1.1) # Measured cold on all four ranks: 7.130 GB (7.060 GB in production), all # but the boundaries a transient recompute workspace; this raw floor # (8.43 GB) covers it. Today's cold admission was 4.637 GB. @@ -216,9 +218,9 @@ def test_gdn_segment_states_are_priced_with_the_segments(): r = tp_rank() rows = 8192 cost = _required(r, group_rows=((rows, True),), gdn_segments=4096) - assert cost.checkpoint_workspace == (4096 + 3) * SEGMENT + assert cost.checkpoint_workspace == (4096 + 3) * SEGMENT + COLD assert cost.required == int( - (OUTPUT + 2 * rows // 4 * LAYERS * H * 2 + (4096 + 3) * SEGMENT) * 1.1 + (OUTPUT + 2 * rows // 4 * LAYERS * H * 2 + (4096 + 3) * SEGMENT + COLD) * 1.1 ) @@ -227,7 +229,7 @@ def test_tp_padding_roots_carry_their_own_states(): r = tp_rank() cost = _required(r, group_rows=((4, True),), gdn_segments=1) # Four roots' initial states alone: 4 x 12 value heads x 128 x 128 x fp32. - assert cost.checkpoint_workspace == 4 * SEGMENT > 4 * 12 * 128 * 128 * 4 + assert cost.checkpoint_workspace - COLD == 4 * SEGMENT > 4 * 12 * 128 * 128 * 4 assert cost.required > 4 * SEGMENT From ecf2ce21bc6fcbb657cbcf9cd6044dfdb1f63821 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 02:35:34 +0000 Subject: [PATCH 15/40] Count head-tied and custom checkpoint gradients as live throughout Name gradient slots with sorted kind/name JSON instead of a hash. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 53 +++++++++++++------ ...st_trainer_rank_adapter_gradient_memory.py | 29 ++++++++-- .../unit/test_trainer_rank_planner_reports.py | 2 +- 3 files changed, 62 insertions(+), 22 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 01d2953c5..78dfbc76b 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -20,6 +20,7 @@ from dataclasses import field as dataclass_field from functools import partial import hashlib +import json import logging import math import os @@ -1068,10 +1069,10 @@ class _SubforwardCost: # Adapter gradients the recompute backward holds beyond the boundaries it # has released (``_checkpoint_adapter_gradient_bytes``), before the safety # factor. Split children training the same slots share them, so a split - # charges the largest once; ``..._slots`` identifies those slots within - # this process (a hash, 0 when none), keeping the cost JSON-serializable. + # charges the largest once; ``..._slots`` names those slots (sorted JSON + # of kind/name pairs, "" when none), keeping the cost JSON-serializable. checkpoint_adapter_gradient: int = 0 - checkpoint_adapter_gradient_slots: int = 0 + checkpoint_adapter_gradient_slots: str = "" @property def ephemeral(self) -> int: @@ -4273,21 +4274,38 @@ def _pending_adapter_gradient_bytes( layer_of: dict[int, int] = {} params: dict[int, torch.nn.Parameter] = {} - def slot_params(module: torch.nn.Module) -> Iterator[torch.nn.Parameter]: - for child in module.modules(): + def slot_params( + modules: Iterable[torch.nn.Module], + ) -> Iterator[torch.nn.Parameter]: + for module in modules: # A LoRA without slot tables holds no slot parameters. - if isinstance(child, LoRA) and "_slot_keys" in vars(child): + if isinstance(module, LoRA) and "_slot_keys" in vars(module): for ref in refs: - yield from child.lora_slot_params(ref) + yield from module.lora_slot_params(ref) for index, layer in enumerate(layers): - for param in slot_params(layer): + for param in slot_params(layer.modules()): params[id(param)] = param layer_of[id(param)] = max(layer_of.get(id(param), -1), index) - for param in slot_params(chunk): - if id(param) not in params: - params[id(param)] = param - layer_of[id(param)] = len(layers) + # Outside the decoder (a head runs its backward first) is live + # throughout, even for a parameter a decoder layer also uses. + inside = {id(module) for module in layers.modules()} + outside = (module for module in chunk.modules() if id(module) not in inside) + for param in slot_params(outside): + params[id(param)] = param + layer_of[id(param)] = len(layers) + # A checkpoint's other trainable parameters (custom objects) have no + # decoder position; count them as live throughout. + for ref in refs: + slot = ( + None + if ref.name is None + else getattr(self, "_checkpoint_slots", {}).get(ref.name) + ) + for param in () if slot is None else slot.params: + if id(param) not in params: + params[id(param)] = param + layer_of[id(param)] = len(layers) sizes = [0] * (len(layers) + 1) for param_id, param in params.items(): if ( @@ -4308,8 +4326,9 @@ def _checkpoint_adapter_gradient_bytes( gradient allocated so far: those of layers i..L-1 (a layer allocates its own during its backward) and any outside the decoder. The floor already prices all L boundaries at once, so the extra peak is - max(0, max over i of gradients(i..) - boundaries(i+1..)). It is taken - over the real per-layer sizes, not a uniform-layer line. A short + max(0, max over i of gradients(i..) - boundaries(i+1..)), taken at every + layer over the real per-layer gradient sizes and the caller's per-layer + boundaries, not along a uniform-layer line. A short first wave peaks at layer 0 (Qwen3.6-35B-A3B CP2: 830-900 MB of expert LoRA gradients live at its peak), a long one at the last layer. ``boundaries`` gives each decoder layer's saved-boundary bytes; @@ -4435,9 +4454,11 @@ def _subforward_cost( checkpoint_peak_increment=required - forward_required, hybridep_growth=hybridep_growth_bytes, checkpoint_adapter_gradient=adapter_gradient, - checkpoint_adapter_gradient_slots=hash(gradient_slots) + checkpoint_adapter_gradient_slots=json.dumps( + sorted([ref.kind, ref.name] for ref in gradient_slots) + ) if adapter_gradient - else 0, + else "", ) def _retained_memory_bytes( diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py index dad274d7d..34b686ccc 100644 --- a/tests/unit/test_trainer_rank_adapter_gradient_memory.py +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -142,6 +142,25 @@ def test_pending_gradients_count_local_unallocated_slot_parameters(): assert r._pending_adapter_gradient_bytes([]) == () +def test_a_parameter_the_head_also_uses_is_live_throughout(): + tied = parameter(7) + layers = [lora(policy=[tied, parameter(3)]), lora(policy=[parameter(5)])] + r = adapter_rank(layers, lora(policy=[tied])) + # The head's backward runs before the decoder's and allocates it first. + assert r._pending_adapter_gradient_bytes([POLICY]) == (3 * 2, 5 * 2, 7 * 2) + + +def test_other_checkpoint_parameters_are_live_throughout(): + from art.trainer_rank._impl import _CheckpointSlot + + adapter = parameter(3) + custom = parameter(11) + r = adapter_rank([lora(policy=[adapter])], torch.nn.Module()) + r._checkpoint_slots["policy"] = _CheckpointSlot(params=(adapter, custom)) + # A custom object's parameter has no decoder position; LoRA ones keep theirs. + assert r._pending_adapter_gradient_bytes([POLICY]) == (3 * 2, 11 * 2) + + def test_a_step_with_allocated_gradients_prices_no_extra(): params = [parameter(100) for _ in range(4)] r = adapter_rank([lora(policy=[p]) for p in params], torch.nn.Module()) @@ -184,7 +203,7 @@ def test_cost_and_estimate_charge_the_extra_while_gradients_are_pending(monkeypa assert extra > 0 cost = priced(r, values, (POLICY, None)) assert cost.checkpoint_adapter_gradient == extra - assert cost.checkpoint_adapter_gradient_slots == hash(frozenset({POLICY})) + assert cost.checkpoint_adapter_gradient_slots == '[["checkpoint", "policy"]]' assert cost.required == int( (out + retained + workspace + COLD + retained + extra) * 1.1 ) @@ -212,7 +231,7 @@ def test_cost_and_estimate_charge_the_extra_while_gradients_are_pending(monkeypa ) -def child(extra: int, slots: int, workspace: int = 10) -> _SubforwardCost: +def child(extra: int, slots: str, workspace: int = 10) -> _SubforwardCost: return _SubforwardCost( required=int((1 + workspace + 1 + extra) * 1.1), retained=0, @@ -225,12 +244,12 @@ def child(extra: int, slots: int, workspace: int = 10) -> _SubforwardCost: def test_split_charges_shared_gradients_once_and_distinct_slots_each(): - shared = [child(500, 7), child(300, 7)] + shared = [child(500, "a"), child(300, "a")] assert TrainerRank._split_required_memory(shared) == int((2 + 2 + 10 + 500) * 1.1) - distinct = [child(500, 7), child(300, 9)] + distinct = [child(500, "a"), child(300, "b")] assert TrainerRank._split_required_memory(distinct) == int((2 + 2 + 10 + 800) * 1.1) # A child with no pending gradients does not change the shared charge. - assert TrainerRank._split_required_memory([child(500, 7), child(0, 0)]) == int( + assert TrainerRank._split_required_memory([child(500, "a"), child(0, "")]) == int( (2 + 2 + 10 + 500) * 1.1 ) diff --git a/tests/unit/test_trainer_rank_planner_reports.py b/tests/unit/test_trainer_rank_planner_reports.py index 22bddb218..b5677bf28 100644 --- a/tests/unit/test_trainer_rank_planner_reports.py +++ b/tests/unit/test_trainer_rank_planner_reports.py @@ -299,7 +299,7 @@ def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): "checkpoint_peak_increment": 0, "hybridep_growth": 0, "checkpoint_adapter_gradient": 0, - "checkpoint_adapter_gradient_slots": 0, + "checkpoint_adapter_gradient_slots": "", }, } ], From 067048eb5e5549239b5c74106189992f7bf5bfba Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 02:54:38 +0000 Subject: [PATCH 16/40] Order slot names safely and leave base-model groups out of adapter slots Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 63 ++++++++++++------- ...st_trainer_rank_adapter_gradient_memory.py | 25 ++++++++ 2 files changed, 67 insertions(+), 21 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 78dfbc76b..073991611 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4245,6 +4245,20 @@ def _checkpoint_memory_floor( workspace = max(workspace, -(-rows // 4) * 4 * self._hidden_size * 2) return retained, workspace + @staticmethod + def _gradient_slots( + group_rows: Sequence[tuple[int, bool]], + slot_refs: Sequence["LoRASlotRef | None"] | None, + ) -> frozenset["LoRASlotRef"]: + """Gradient groups' adapter slots; the base model (no name) has none.""" + return frozenset( + ref + for (_, grad), ref in zip( + group_rows, slot_refs or (None,) * len(group_rows), strict=True + ) + if grad and ref is not None and ref.name is not None + ) + def _pending_adapter_gradient_bytes( self, refs: Iterable["LoRASlotRef"] ) -> tuple[int, ...]: @@ -4288,9 +4302,18 @@ def slot_params( params[id(param)] = param layer_of[id(param)] = max(layer_of.get(id(param), -1), index) # Outside the decoder (a head runs its backward first) is live - # throughout, even for a parameter a decoder layer also uses. - inside = {id(module) for module in layers.modules()} - outside = (module for module in chunk.modules() if id(module) not in inside) + # throughout, even for a parameter or module a decoder layer also uses: + # walk every path except through the layers themselves. + outside: list[torch.nn.Module] = [] + visited: set[int] = set() + pending_modules: list[torch.nn.Module] = [chunk] + while pending_modules: + module = pending_modules.pop() + if module is layers or id(module) in visited: + continue + visited.add(id(module)) + outside.append(module) + pending_modules.extend(module.children()) for param in slot_params(outside): params[id(param)] = param layer_of[id(param)] = len(layers) @@ -4407,13 +4430,7 @@ def _subforward_cost( # backing stores, nor a bound for compiler saves or other backward work. # Keep it out of forward retention, including the cold fallback above. gradient = checkpoint_retained - gradient_slots = frozenset( - ref - for (_, grad), ref in zip( - group_rows, slot_refs or (None,) * len(group_rows), strict=True - ) - if grad and ref is not None - ) + gradient_slots = self._gradient_slots(group_rows, slot_refs) adapter_gradient = ( self._checkpoint_adapter_gradient_bytes( gradient_slots, (gradient // self._num_layers,) * self._num_layers @@ -4455,7 +4472,17 @@ def _subforward_cost( hybridep_growth=hybridep_growth_bytes, checkpoint_adapter_gradient=adapter_gradient, checkpoint_adapter_gradient_slots=json.dumps( - sorted([ref.kind, ref.name] for ref in gradient_slots) + [ + [ref.kind, ref.name] + for ref in sorted( + gradient_slots, + key=lambda ref: ( + ref.kind, + ref.name is not None, + ref.name or "", + ), + ) + ] ) if adapter_gradient else "", @@ -6220,7 +6247,9 @@ def _estimate_flat_forward( # exact plan instead of admitting with the constructor rank. return None gradient_slots = [ - ref for (ref, grad), _ in groups if grad and ref is not None + ref + for (ref, grad), _ in groups + if grad and ref is not None and ref.name is not None ] if ( gradient_slots @@ -8245,15 +8274,7 @@ def _estimate_required_memory_bytes_from_values( if include_checkpoint_input_gradient and retained: # The backward's other end and cold transients, as _subforward_cost. backward = retained + self._checkpoint_adapter_gradient_bytes( - ( - ref - for (_, grad), ref in zip( - group_rows, - slot_refs or (None,) * len(group_rows), - strict=True, - ) - if grad and ref is not None - ), + self._gradient_slots(group_rows, slot_refs), (retained // self._num_layers,) * self._num_layers, ) if profiled is None and any(grad for _, grad in group_rows): diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py index 34b686ccc..c3aa0ae79 100644 --- a/tests/unit/test_trainer_rank_adapter_gradient_memory.py +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -150,6 +150,31 @@ def test_a_parameter_the_head_also_uses_is_live_throughout(): assert r._pending_adapter_gradient_bytes([POLICY]) == (3 * 2, 5 * 2, 7 * 2) +def test_a_module_the_head_also_uses_is_live_throughout(): + shared = lora(policy=[parameter(7)]) + r = adapter_rank([shared, torch.nn.Module()], shared) + assert r._pending_adapter_gradient_bytes([POLICY]) == (0, 0, 7 * 2) + + +def test_base_model_groups_own_no_adapter_gradients(monkeypatch): + r = rank() + values = r._estimate_flat_forward(requests(67, 4096)) + with_pending(monkeypatch, r, [23 * 2**20] * 40 + [0]) + n, out, signature, groups, head = values + both = r._subforward_cost( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=((33, True), (34, True), (4096, False)), + slot_refs=(POLICY, LoRASlotRef("checkpoint", None), None), + head_workspace_bytes=head, + ) + # The base group's slot has no adapter, so a split with or without it + # still shares one slot's gradients. + assert both.checkpoint_adapter_gradient_slots == '[["checkpoint", "policy"]]' + + def test_other_checkpoint_parameters_are_live_throughout(): from art.trainer_rank._impl import _CheckpointSlot From c35b238798076e6c34e02d87823b32cc74184e01 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 06:33:29 +0000 Subject: [PATCH 17/40] Test that base-model gradient groups price and defer nothing Co-Authored-By: Claude Opus 5.5 (1M context) --- ...st_trainer_rank_adapter_gradient_memory.py | 34 +++++++++++++++++++ 1 file changed, 34 insertions(+) diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py index c3aa0ae79..627499d3a 100644 --- a/tests/unit/test_trainer_rank_adapter_gradient_memory.py +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -300,3 +300,37 @@ def pending(refs): assert seen == [(POLICY,)] monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", lambda refs: ()) assert r._estimate_flat_forward(requests(67, 4096)) is not None + + +def test_a_base_only_gradient_group_prices_and_defers_nothing(monkeypatch): + r = rank() + values = r._estimate_flat_forward(requests(67, 4096)) + n, out, signature, groups, head = values + base = LoRASlotRef("checkpoint", None) + with_pending(monkeypatch, r, [23 * 2**20] * 40 + [0]) + cost = priced(r, values, (base, None)) + # The base model has no adapter: no extra, no slot identity. + assert cost.checkpoint_adapter_gradient == 0 + assert cost.checkpoint_adapter_gradient_slots == "" + estimate = r._estimate_required_memory_bytes_from_values( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=(base, None), + head_workspace_bytes=head, + ) + assert estimate == cost.required + # Nor does the cheap estimate defer to the exact plan for it. + monkeypatch.setattr(r, "_ensure_checkpoint_slots_for", lambda *a, **k: None) + monkeypatch.setattr(r, "_resolve_slot_ref", lambda request, checkpoint: base) + seen: list[tuple[LoRASlotRef, ...]] = [] + + def pending(refs): + seen.append(tuple(refs)) + return (1,) * 41 + + monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", pending) + assert r._estimate_flat_forward(requests(67, 4096)) is not None + assert seen == [] From e54f4bc25814f9d547145721b5695661d5308c54 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:08:04 +0000 Subject: [PATCH 18/40] Walk gradient groups' backward one after another, in the worst order Autograd drains the last-forwarded group's chain before an earlier one's, and separate backward calls may come in either order, so a short group's gradients can peak beside another group's unreleased boundaries. Price each gradient group against its own boundaries, with groups not yet run holding theirs and groups already run holding their gradients, over every order. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 73 ++++++++--- ...st_trainer_rank_adapter_gradient_memory.py | 114 +++++++++++++++++- 2 files changed, 167 insertions(+), 20 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 073991611..46750c08d 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -20,6 +20,7 @@ from dataclasses import field as dataclass_field from functools import partial import hashlib +import itertools import json import logging import math @@ -4339,8 +4340,31 @@ def slot_params( sizes[layer_of[param_id]] += param.numel() * param.element_size() return tuple(sizes) if any(sizes) else () + def _checkpoint_gradient_groups( + self, + group_rows: Sequence[tuple[int, bool]], + slot_refs: Sequence["LoRASlotRef | None"] | None, + ) -> tuple[tuple["LoRASlotRef | None", tuple[int, ...]], ...]: + """Each gradient group's adapter slot and per-layer saved boundaries. + + In execution order, as ``_checkpoint_memory_floor`` prices them: every + decoder layer saves the group's rows (this rank's TP shard). The base + model (no name) has no adapter slot. + """ + tp = self._topology_key()[1] + return tuple( + ( + ref if ref is not None and ref.name is not None else None, + (-(-rows // tp) * self._hidden_size * 2,) * self._num_layers, + ) + for (rows, grad), ref in zip( + group_rows, slot_refs or (None,) * len(group_rows), strict=True + ) + if grad + ) + def _checkpoint_adapter_gradient_bytes( - self, slots: Iterable["LoRASlotRef"], boundaries: Sequence[int] + self, groups: Sequence[tuple["LoRASlotRef | None", Sequence[int]]] ) -> int: """The recompute backward's adapter-gradient peak beyond released boundaries. @@ -4354,19 +4378,39 @@ def _checkpoint_adapter_gradient_bytes( boundaries, not along a uniform-layer line. A short first wave peaks at layer 0 (Qwen3.6-35B-A3B CP2: 830-900 MB of expert LoRA gradients live at its peak), a long one at the last layer. - ``boundaries`` gives each decoder layer's saved-boundary bytes; - ``slots`` are the gradient groups' adapter slots. + ``groups`` gives each gradient group's adapter slot (None for the base + model) and each decoder layer's saved-boundary bytes. Groups run their + backward one after another, not layer by layer together: autograd + drains the last-forwarded group's chain first, and separate backward + calls can come in either order. While one group runs, a group yet to + run still holds all its boundaries and one already run all its + gradients, so take the worst order. """ - pending = self._pending_adapter_gradient_bytes(slots) - if not pending or len(pending) != len(boundaries) + 1: + chains = [] + for slot, boundaries in groups: + pending = ( + () if slot is None else self._pending_adapter_gradient_bytes((slot,)) + ) + if pending and len(pending) != len(boundaries) + 1: + return 0 + chains.append((pending or (0,) * (len(boundaries) + 1), boundaries)) + if not any(any(pending) for pending, _ in chains): return 0 - extra = gradients = pending[-1] - released = 0 - for index in range(len(boundaries) - 1, -1, -1): - gradients += pending[index] - extra = max(extra, gradients - released) - released += boundaries[index] - return extra + if len(chains) > 4: + # Too many orders to walk: every gradient live, nothing released. + return sum(sum(pending) for pending, _ in chains) + worst = 0 + for order in itertools.permutations(chains): + allocated = released = 0 + for pending, boundaries in order: + gradients = allocated + pending[-1] + worst = max(worst, gradients - released) + for index in range(len(boundaries) - 1, -1, -1): + gradients += pending[index] + worst = max(worst, gradients - released) + released += boundaries[index] + allocated = gradients + return worst def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( @@ -4433,7 +4477,7 @@ def _subforward_cost( gradient_slots = self._gradient_slots(group_rows, slot_refs) adapter_gradient = ( self._checkpoint_adapter_gradient_bytes( - gradient_slots, (gradient // self._num_layers,) * self._num_layers + self._checkpoint_gradient_groups(group_rows, slot_refs) ) if gradient else 0 @@ -8274,8 +8318,7 @@ def _estimate_required_memory_bytes_from_values( if include_checkpoint_input_gradient and retained: # The backward's other end and cold transients, as _subforward_cost. backward = retained + self._checkpoint_adapter_gradient_bytes( - self._gradient_slots(group_rows, slot_refs), - (retained // self._num_layers,) * self._num_layers, + self._checkpoint_gradient_groups(group_rows, slot_refs) ) if profiled is None and any(grad for _, grad in group_rows): backward += _COLD_RECOMPUTE_TRANSIENT_BYTES diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py index 627499d3a..0e80b9e2e 100644 --- a/tests/unit/test_trainer_rank_adapter_gradient_memory.py +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -5,6 +5,7 @@ """ from collections.abc import Sequence +import itertools import random import pytest @@ -61,7 +62,7 @@ def test_extra_is_the_worst_backward_layer_over_real_sizes( ): r = rank() with_pending(monkeypatch, r, pending) - assert r._checkpoint_adapter_gradient_bytes((POLICY,), boundaries) == oracle( + assert r._checkpoint_adapter_gradient_bytes(((POLICY, boundaries),)) == oracle( pending, boundaries ) @@ -76,14 +77,103 @@ def test_extra_matches_every_layer_for_random_sizes(monkeypatch): ] boundaries = [generator.randint(0, 40) for _ in range(layers)] with_pending(monkeypatch, r, pending) - extra = r._checkpoint_adapter_gradient_bytes((POLICY,), boundaries) + extra = r._checkpoint_adapter_gradient_bytes(((POLICY, boundaries),)) assert extra == oracle(pending, boundaries) def test_no_pending_gradients_add_nothing(monkeypatch): r = rank() with_pending(monkeypatch, r, ()) - assert r._checkpoint_adapter_gradient_bytes((POLICY,), [4] * 40) == 0 + assert r._checkpoint_adapter_gradient_bytes(((POLICY, [4] * 40),)) == 0 + + +def sequential_oracle(chains): + """Worst live bytes beyond the floor, one group's backward after another. + + While a group recomputes layer i, groups run before it hold all their + gradients and none of their boundaries; groups yet to run hold all their + boundaries; the running group releases its boundaries above i. + """ + worst = 0 + for order in itertools.permutations(chains): + for position, (pending, boundaries) in enumerate(order): + done = order[:position] + layers = len(boundaries) + for index in range(layers): + gradients = ( + sum(sum(p) for p, _ in done) + + pending[layers] + + sum(pending[index:layers]) + ) + released = sum(sum(b) for _, b in done) + sum(boundaries[index + 1 :]) + worst = max(worst, gradients - released) + return worst + + +def with_slot_pending(monkeypatch, r, by_slot): + monkeypatch.setattr( + r, + "_pending_adapter_gradient_bytes", + lambda refs: by_slot.get(tuple(refs), ()), + ) + + +def test_gradient_groups_run_their_backward_one_after_another(monkeypatch): + r = rank() + # A short policy group beside a long group of another slot: whichever + # runs first, the other keeps its boundaries (or its gradients) meanwhile. + policy = ([30] * 4 + [0], [2] * 4) + other = ([5] * 4 + [1], [40] * 4) + with_slot_pending(monkeypatch, r, {(POLICY,): policy[0], (OTHER,): other[0]}) + extra = r._checkpoint_adapter_gradient_bytes( + ((POLICY, policy[1]), (OTHER, other[1])) + ) + assert extra == sequential_oracle([policy, other]) + # One chain with every group's boundaries released together would have + # priced far less. + combined = [a + b for a, b in zip(policy[0], other[0])] + assert extra > oracle(combined, [a + b for a, b in zip(policy[1], other[1])]) + # A base-model group owns no gradients, but its boundaries stay live + # while the policy group runs first. + base = ([0] * 5, [40] * 4) + assert ( + r._checkpoint_adapter_gradient_bytes(((POLICY, policy[1]), (None, base[1]))) + == sequential_oracle([policy, base]) + == oracle(*policy) + ) + + +def test_sequential_groups_match_every_order_for_random_sizes(monkeypatch): + r = rank() + generator = random.Random(1) + slots = [POLICY, OTHER, LoRASlotRef("checkpoint", "third")] + for _ in range(200): + layers = generator.randint(1, 6) + chains, groups, by_slot = [], [], {} + for slot in slots[: generator.randint(1, 3)]: + pending = [generator.randint(0, 30) for _ in range(layers + 1)] + boundaries = [generator.randint(0, 30) for _ in range(layers)] + if generator.random() < 0.25: + slot, pending = None, [0] * (layers + 1) + else: + by_slot[(slot,)] = pending + chains.append((pending, boundaries)) + groups.append((slot, boundaries)) + with_slot_pending(monkeypatch, r, by_slot) + assert r._checkpoint_adapter_gradient_bytes(groups) == sequential_oracle(chains) + + +def test_many_gradient_groups_price_every_gradient_live(monkeypatch): + r = rank() + slots = [LoRASlotRef("checkpoint", f"slot{index}") for index in range(5)] + pending = [3, 1, 2] + with_slot_pending(monkeypatch, r, {(slot,): pending for slot in slots}) + groups = [(slot, [100, 100]) for slot in slots] + # Too many orders to walk: all five slots' gradients, nothing released. + assert r._checkpoint_adapter_gradient_bytes(groups) == 5 * sum(pending) + assert r._checkpoint_adapter_gradient_bytes(groups[:4]) == sequential_oracle( + [(pending, [100, 100])] * 4 + ) def lora(**slots: list[torch.nn.Parameter]) -> LoRA: @@ -159,9 +249,10 @@ def test_a_module_the_head_also_uses_is_live_throughout(): def test_base_model_groups_own_no_adapter_gradients(monkeypatch): r = rank() values = r._estimate_flat_forward(requests(67, 4096)) - with_pending(monkeypatch, r, [23 * 2**20] * 40 + [0]) + pending = [23 * 2**20] * 40 + [0] + with_pending(monkeypatch, r, pending) n, out, signature, groups, head = values - both = r._subforward_cost( + values = dict( packed_tokens=n, output_bytes=out, signature=signature, @@ -170,9 +261,22 @@ def test_base_model_groups_own_no_adapter_gradients(monkeypatch): slot_refs=(POLICY, LoRASlotRef("checkpoint", None), None), head_workspace_bytes=head, ) + both = r._subforward_cost(**values) # The base group's slot has no adapter, so a split with or without it # still shares one slot's gradients. assert both.checkpoint_adapter_gradient_slots == '[["checkpoint", "policy"]]' + # But its boundaries stay live while the policy group's backward runs. + boundary = 2048 * 2 * 40 + expected = sequential_oracle( + [(pending, [33 * 2048 * 2] * 40), ([0] * 41, [34 * 2048 * 2] * 40)] + ) + assert ( + both.checkpoint_adapter_gradient + == expected + == oracle(pending, [33 * 2048 * 2] * 40) + ) + assert expected > oracle(pending, [(33 + 34) * boundary // 40] * 40) + assert r._estimate_required_memory_bytes_from_values(**values) == both.required def test_other_checkpoint_parameters_are_live_throughout(): From 0df629cf516a6237afb77a73fa60ea111c2b398c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:29:36 +0000 Subject: [PATCH 19/40] Price gradient groups' worst backward order in closed form Any set of the other groups can have run first, so the worst order adds every other group whose gradients outweigh its boundaries to one group's own walk. Exact for any number of groups, without walking permutations. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 72 ++++++++++--------- ...st_trainer_rank_adapter_gradient_memory.py | 23 +++--- 2 files changed, 52 insertions(+), 43 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 46750c08d..827f7ffa7 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -20,7 +20,6 @@ from dataclasses import field as dataclass_field from functools import partial import hashlib -import itertools import json import logging import math @@ -4368,23 +4367,9 @@ def _checkpoint_adapter_gradient_bytes( ) -> int: """The recompute backward's adapter-gradient peak beyond released boundaries. - Backward recomputes the last layer first. While it recomputes layer i it - still holds the saved boundaries of layers 0..i and every adapter - gradient allocated so far: those of layers i..L-1 (a layer allocates its - own during its backward) and any outside the decoder. The floor already - prices all L boundaries at once, so the extra peak is - max(0, max over i of gradients(i..) - boundaries(i+1..)), taken at every - layer over the real per-layer gradient sizes and the caller's per-layer - boundaries, not along a uniform-layer line. A short - first wave peaks at layer 0 (Qwen3.6-35B-A3B CP2: 830-900 MB of expert - LoRA gradients live at its peak), a long one at the last layer. ``groups`` gives each gradient group's adapter slot (None for the base - model) and each decoder layer's saved-boundary bytes. Groups run their - backward one after another, not layer by layer together: autograd - drains the last-forwarded group's chain first, and separate backward - calls can come in either order. While one group runs, a group yet to - run still holds all its boundaries and one already run all its - gradients, so take the worst order. + model) and each decoder layer's saved-boundary bytes + (``_adapter_gradient_walk``). """ chains = [] for slot, boundaries in groups: @@ -4394,22 +4379,45 @@ def _checkpoint_adapter_gradient_bytes( if pending and len(pending) != len(boundaries) + 1: return 0 chains.append((pending or (0,) * (len(boundaries) + 1), boundaries)) - if not any(any(pending) for pending, _ in chains): - return 0 - if len(chains) > 4: - # Too many orders to walk: every gradient live, nothing released. - return sum(sum(pending) for pending, _ in chains) + return self._adapter_gradient_walk(chains) + + @staticmethod + def _adapter_gradient_walk( + chains: Sequence[tuple[Sequence[int], Sequence[int]]], + ) -> int: + """The adapter-gradient peak beyond the floor over gradient groups' backward. + + Each chain is a gradient group's pending gradient bytes (per decoder + layer, then outside the decoder) and saved-boundary bytes per layer. + Backward recomputes the last layer first. While it recomputes layer i it + still holds the saved boundaries of layers 0..i and every adapter + gradient allocated so far: those of layers i..L-1 (a layer allocates its + own during its backward) and any outside the decoder. The floor already + prices all L boundaries at once, so one group's extra peak is + max(0, max over i of gradients(i..) - boundaries(i+1..)), taken at every + layer over the real per-layer gradient sizes and the caller's per-layer + boundaries, not along a uniform-layer line. A short + first wave peaks at layer 0 (Qwen3.6-35B-A3B CP2: 830-900 MB of expert + LoRA gradients live at its peak), a long one at the last layer. + Groups run their backward one after another, not layer by layer + together: autograd drains the last-forwarded group's chain first, and + separate backward calls can come in either order. While one group runs, + each group already run holds all its gradients and none of its + boundaries, and each group yet to run all its boundaries. Any set of the + other groups can have run first, so the worst adds every other group + whose gradients outweigh its boundaries. + """ + nets = [sum(pending) - sum(boundaries) for pending, boundaries in chains] + others = sum(max(0, net) for net in nets) worst = 0 - for order in itertools.permutations(chains): - allocated = released = 0 - for pending, boundaries in order: - gradients = allocated + pending[-1] - worst = max(worst, gradients - released) - for index in range(len(boundaries) - 1, -1, -1): - gradients += pending[index] - worst = max(worst, gradients - released) - released += boundaries[index] - allocated = gradients + for (pending, boundaries), net in zip(chains, nets, strict=True): + extra = gradients = pending[-1] + released = 0 + for index in range(len(boundaries) - 1, -1, -1): + gradients += pending[index] + extra = max(extra, gradients - released) + released += boundaries[index] + worst = max(worst, extra + others - max(0, net)) return worst def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py index 0e80b9e2e..71721224a 100644 --- a/tests/unit/test_trainer_rank_adapter_gradient_memory.py +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -146,11 +146,11 @@ def test_gradient_groups_run_their_backward_one_after_another(monkeypatch): def test_sequential_groups_match_every_order_for_random_sizes(monkeypatch): r = rank() generator = random.Random(1) - slots = [POLICY, OTHER, LoRASlotRef("checkpoint", "third")] + slots = [LoRASlotRef("checkpoint", f"slot{index}") for index in range(5)] for _ in range(200): layers = generator.randint(1, 6) chains, groups, by_slot = [], [], {} - for slot in slots[: generator.randint(1, 3)]: + for slot in slots[: generator.randint(1, 5)]: pending = [generator.randint(0, 30) for _ in range(layers + 1)] boundaries = [generator.randint(0, 30) for _ in range(layers)] if generator.random() < 0.25: @@ -163,17 +163,18 @@ def test_sequential_groups_match_every_order_for_random_sizes(monkeypatch): assert r._checkpoint_adapter_gradient_bytes(groups) == sequential_oracle(chains) -def test_many_gradient_groups_price_every_gradient_live(monkeypatch): +def test_many_gradient_groups_price_exactly(monkeypatch): r = rank() - slots = [LoRASlotRef("checkpoint", f"slot{index}") for index in range(5)] - pending = [3, 1, 2] - with_slot_pending(monkeypatch, r, {(slot,): pending for slot in slots}) - groups = [(slot, [100, 100]) for slot in slots] - # Too many orders to walk: all five slots' gradients, nothing released. - assert r._checkpoint_adapter_gradient_bytes(groups) == 5 * sum(pending) - assert r._checkpoint_adapter_gradient_bytes(groups[:4]) == sequential_oracle( - [(pending, [100, 100])] * 4 + slots = [LoRASlotRef("checkpoint", f"slot{index}") for index in range(6)] + # Only groups whose gradients outweigh their boundaries raise another + # group's peak by having run first. + chains = [([3, 1, 2], [100, 100])] * 3 + [([90, 40, 5], [2, 1])] * 3 + with_slot_pending( + monkeypatch, r, {(slot,): pending for slot, (pending, _) in zip(slots, chains)} ) + groups = [(slot, boundaries) for slot, (_, boundaries) in zip(slots, chains)] + assert r._checkpoint_adapter_gradient_bytes(groups) == sequential_oracle(chains) + assert sequential_oracle(chains) < sum(sum(pending) for pending, _ in chains) def lora(**slots: list[torch.nn.Parameter]) -> LoRA: From 2a6419528d8f4071d03e80c2983bd062746345ee Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:44:22 +0000 Subject: [PATCH 20/40] Test that a layer-count mismatch prices no adapter-gradient extra Co-Authored-By: Claude Opus 5.5 (1M context) --- .../test_trainer_rank_adapter_gradient_memory.py | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py index 71721224a..0da30ce0b 100644 --- a/tests/unit/test_trainer_rank_adapter_gradient_memory.py +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -87,6 +87,20 @@ def test_no_pending_gradients_add_nothing(monkeypatch): assert r._checkpoint_adapter_gradient_bytes(((POLICY, [4] * 40),)) == 0 +def test_a_layer_count_mismatch_prices_no_extra(monkeypatch): + r = rank() + with_slot_pending( + monkeypatch, r, {(POLICY,): [5] * 4 + [0], (OTHER,): [9] * 3 + [0]} + ) + # Pending gradients and boundaries must describe the same decoder layers; + # otherwise the whole term is left out rather than misplaced. + assert r._checkpoint_adapter_gradient_bytes(((POLICY, [1] * 4),)) > 0 + assert r._checkpoint_adapter_gradient_bytes(((OTHER, [1] * 4),)) == 0 + assert ( + r._checkpoint_adapter_gradient_bytes(((POLICY, [1] * 4), (OTHER, [1] * 4))) == 0 + ) + + def sequential_oracle(chains): """Worst live bytes beyond the floor, one group's backward after another. From 23f4f9582f3b60114233ef4d18d7b5f8fc9370a7 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 13:17:52 +0000 Subject: [PATCH 21/40] Price the head and decoder recompute backward stages apart The head finishes its backward before the decoder's recompute starts, so its buffers never meet a recomputed layer's workspace or the adapter gradients that recompute allocates. Where the MoE stage covers the recompute, price max(head stage, decoder stage) instead of adding the adapter extra on top of the larger of the two workspaces. The head stage holds the final decoder outputs, the checkpointed head chunks' selected rows and the hidden-row gradient (each at most the gradient rows x 2H), the head buffers, TE's first-GEMM workspaces and other groups' adapter gradients. A split keeps each child's whole head stage beside the children's adapter gradients. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 144 ++++++++++--- ...trainer_rank_checkpoint_gradient_memory.py | 6 +- tests/unit/test_trainer_rank_head_memory.py | 13 +- .../test_trainer_rank_head_stage_memory.py | 197 ++++++++++++++++++ 4 files changed, 321 insertions(+), 39 deletions(-) create mode 100644 tests/unit/test_trainer_rank_head_stage_memory.py diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 55b7741a5..0b9c29789 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4534,9 +4534,20 @@ def _checkpoint_input_gradient_bytes( retained, _ = self._checkpoint_memory_floor(group_rows) if not retained: return 0 - gradient_rows = sum(rows for rows, grad in group_rows if grad) + if self._checkpoint_gradient_covered(group_rows, slot_refs): + return ( + sum(rows for rows, grad in group_rows if grad) * self._hidden_size * 2 + ) + return retained + + def _checkpoint_gradient_covered( + self, + group_rows: tuple[tuple[int, bool], ...], + slot_refs: tuple["LoRASlotRef | None", ...] | None, + ) -> bool: + """Whether the traced MoE stage covers every gradient group's recompute.""" refs = (None,) * len(group_rows) if slot_refs is None else slot_refs - if ( + return bool( self._topology_key()[2] <= 2 and self._checkpoint_moe_bytes_per_token() and all( @@ -4544,9 +4555,7 @@ def _checkpoint_input_gradient_bytes( for (_, grad), ref in zip(group_rows, refs, strict=True) if grad ) - ): - return gradient_rows * self._hidden_size * 2 - return retained + ) @staticmethod def _gradient_slots( @@ -4666,13 +4675,17 @@ def _checkpoint_gradient_groups( ) def _checkpoint_adapter_gradient_bytes( - self, groups: Sequence[tuple["LoRASlotRef | None", Sequence[int]]] + self, + groups: Sequence[tuple["LoRASlotRef | None", Sequence[int]]], + *, + head: bool = False, ) -> int: """The recompute backward's adapter-gradient peak beyond released boundaries. ``groups`` gives each gradient group's adapter slot (None for the base model) and each decoder layer's saved-boundary bytes - (``_adapter_gradient_walk``). + (``_adapter_gradient_walk``). With ``head``, those live while a group's + head runs its backward instead (``_adapter_gradient_head``). """ chains = [] for slot, boundaries in groups: @@ -4682,8 +4695,33 @@ def _checkpoint_adapter_gradient_bytes( if pending and len(pending) != len(boundaries) + 1: return 0 chains.append((pending or (0,) * (len(boundaries) + 1), boundaries)) + if head: + return self._adapter_gradient_head(chains) return self._adapter_gradient_walk(chains) + @staticmethod + def _adapter_gradient_head( + chains: Sequence[tuple[Sequence[int], Sequence[int]]], + ) -> int: + """Adapter gradients live while one group's head runs its backward. + + Chains as ``_adapter_gradient_walk``. None of the group's decoder layers + has run yet; its gradients outside the decoder are live, and so is every + other group whose gradients outweigh its boundaries, as any may have run + first. + """ + nets = [ + max(0, sum(pending) - sum(boundaries)) for pending, boundaries in chains + ] + others = sum(nets) + return max( + ( + pending[-1] + others - net + for (pending, _), net in zip(chains, nets, strict=True) + ), + default=0, + ) + @staticmethod def _adapter_gradient_walk( chains: Sequence[tuple[Sequence[int], Sequence[int]]], @@ -4723,6 +4761,40 @@ def _adapter_gradient_walk( worst = max(worst, extra + others - max(0, net)) return worst + def _checkpoint_head_stage_bytes( + self, + head_workspace_bytes: int, + gradient: int, + group_rows: tuple[tuple[int, bool], ...], + slot_refs: tuple["LoRASlotRef | None", ...] | None, + ) -> int | None: + """The head backward's peak beyond the boundaries and ``gradient``. + + The head finishes its backward before the decoder's recompute starts, + so its buffers never meet a recomputed layer's workspace or the adapter + gradients that recompute allocates. Where the MoE stage covers the + recompute, ``gradient`` is the gradient groups' rows x 2H, and each of + these is at most that: the final decoder outputs, the hidden rows each + checkpointed head chunk saved (this rank's rows), and the hidden-row + gradient. Beside them are the head's own buffers, TE's cuBLAS workspaces + from the forward's first GEMMs, and the adapter gradients of groups + whose backward ran first. Qwen3.6-35B-A3B CP2 traces (EP1 and EP2) + show these terms at the head's peak. Elsewhere None: the head shares + the decoder stage. + """ + if not head_workspace_bytes or not self._checkpoint_gradient_covered( + group_rows, slot_refs + ): + return None + return ( + head_workspace_bytes + + 2 * gradient + + self._te_workspace_growth_bytes() + + self._checkpoint_adapter_gradient_bytes( + self._checkpoint_gradient_groups(group_rows, slot_refs), head=True + ) + ) + def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( packed_tokens=plan.packed_tokens, @@ -4800,24 +4872,27 @@ def _subforward_cost( checkpoint_retained = output_bytes + max( checkpoint_retained, checkpoint_floor[0] ) - checkpoint_workspace = max( - checkpoint_workspace, head_workspace_bytes, checkpoint_floor[1] - ) - if gradient and self._memory_profiles.get(signature) is None: - checkpoint_workspace += _COLD_RECOMPUTE_TRANSIENT_BYTES + decoder_workspace = max(checkpoint_workspace, checkpoint_floor[1]) + checkpoint_workspace = max(decoder_workspace, head_workspace_bytes) forward_required = required if gradient: + head_stage = self._checkpoint_head_stage_bytes( + head_workspace_bytes, gradient, group_rows, slot_refs + ) + if head_stage is None: + peak = checkpoint_workspace + adapter_gradient + else: + # A split adds its children's adapter gradients to the largest + # workspace, and one child's head can follow another's decoder + # backward: keep the whole head stage there. + checkpoint_workspace = max(decoder_workspace, head_stage) + peak = max(decoder_workspace + adapter_gradient, head_stage) + if self._memory_profiles.get(signature) is None: + checkpoint_workspace += _COLD_RECOMPUTE_TRANSIENT_BYTES + peak += _COLD_RECOMPUTE_TRANSIENT_BYTES required = max( required, - int( - ( - checkpoint_retained - + checkpoint_workspace - + gradient - + adapter_gradient - ) - * _MEMORY_SAFETY_FACTOR - ), + int((checkpoint_retained + gradient + peak) * _MEMORY_SAFETY_FACTOR), ) # HybridEP buffer growth stays allocated through the forward and # backward peaks, but is not forward retention. @@ -8616,23 +8691,26 @@ def _estimate_required_memory_bytes_from_values( gdn_segments, routed_rows=group_routed_rows, ) - backward = 0 + decoder_workspace = max(workspace, checkpoint_floor[1]) + peak = max(decoder_workspace, head_workspace_bytes) if include_checkpoint_input_gradient and retained: # The input gradient, the backward's other end and cold transients, - # as _subforward_cost. - backward = self._checkpoint_input_gradient_bytes( - group_rows, slot_refs - ) + self._checkpoint_adapter_gradient_bytes( + # staged as _subforward_cost. + gradient = self._checkpoint_input_gradient_bytes(group_rows, slot_refs) + adapter_gradient = self._checkpoint_adapter_gradient_bytes( self._checkpoint_gradient_groups(group_rows, slot_refs) ) + head_stage = self._checkpoint_head_stage_bytes( + head_workspace_bytes, gradient, group_rows, slot_refs + ) + peak = gradient + ( + peak + adapter_gradient + if head_stage is None + else max(decoder_workspace + adapter_gradient, head_stage) + ) if profiled is None and any(grad for _, grad in group_rows): - backward += _COLD_RECOMPUTE_TRANSIENT_BYTES - static_compute = max( - static_compute, - max(retained, checkpoint_floor[0]) - + max(workspace, head_workspace_bytes, checkpoint_floor[1]) - + backward, - ) + peak += _COLD_RECOMPUTE_TRANSIENT_BYTES + static_compute = max(static_compute, max(retained, checkpoint_floor[0]) + peak) if signature.topology[2] > 1: # Local head results coexist with full CP outputs during gathering. # Uneven rank plans can assign all of an item's rows to one rank. diff --git a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py index 10b4e82a4..d5bf9e3a3 100644 --- a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py @@ -141,8 +141,10 @@ def test_gradient_is_not_absorbed_by_larger_head_workspace(): cost = price(r, (n, out, sig, groups, head)) boundaries = 67 * 40 * 2048 * 2 gradient = 67 * 2048 * 2 - assert cost.checkpoint_workspace == head + COLD - assert cost.required == int((out + head + COLD + boundaries + gradient) * 1.1) + # The head stage: two more gradient-row terms and TE's first-GEMM workspaces. + stage = head + 2 * gradient + r._te_workspace_growth_bytes() + assert cost.checkpoint_workspace == stage + COLD + assert cost.required == int((out + stage + COLD + boundaries + gradient) * 1.1) assert cost.retained == int((out + head + boundaries) * 1.1) r._memory_profiles[sig] = _MemoryProfile( bytes_per_token=10**9, diff --git a/tests/unit/test_trainer_rank_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py index 71ca1df11..26f44aa7e 100644 --- a/tests/unit/test_trainer_rank_head_memory.py +++ b/tests/unit/test_trainer_rank_head_memory.py @@ -150,9 +150,12 @@ def test_outputs_retention_and_empirical_peak_are_counted_once(): head = 3 * 512 * 248320 * 2 cost = r._plan_cost(plan) assert cost.retained == int((plan.output_bytes + retained + head) * 1.1) - # Unprofiled: the first execution's transients beside the head workspace. + # Unprofiled head stage: the head workspace, the final output, saved + # selected rows and hidden-row gradient beside the incoming one, TE's + # first-GEMM workspaces and the first execution's transients. + te = r._te_workspace_growth_bytes() assert cost.required == int( - (plan.output_bytes + retained + gradient + head + COLD) * 1.1 + (plan.output_bytes + retained + 3 * gradient + head + te + COLD) * 1.1 ) r._memory_profiles[plan.signature] = _MemoryProfile( bytes_per_token=2_000_000, @@ -320,8 +323,10 @@ def test_target_backward_refuses_budget_below_logits_and_both_gradients(rows): retained, _ = r._checkpoint_memory_floor(r._plan_group_rows(plan)) gradient = rows * 2048 * 2 dense = min(rows, 512) * 248320 * 2 - before = int((plan.output_bytes + retained + gradient + 2 * dense + COLD) * 1.1) - expected = int((plan.output_bytes + retained + gradient + 3 * dense + COLD) * 1.1) + before = int((plan.output_bytes + retained + 3 * gradient + 2 * dense + COLD) * 1.1) + expected = int( + (plan.output_bytes + retained + 3 * gradient + 3 * dense + COLD) * 1.1 + ) r._available_memory_bytes = lambda: (before + expected) // 2 check = r._memory_check(plan) assert check.estimated_required_bytes == expected diff --git a/tests/unit/test_trainer_rank_head_stage_memory.py b/tests/unit/test_trainer_rank_head_stage_memory.py new file mode 100644 index 000000000..2d67b68c5 --- /dev/null +++ b/tests/unit/test_trainer_rank_head_stage_memory.py @@ -0,0 +1,197 @@ +"""The head's backward and the decoder's recompute peak apart: CPU admission math. + +Qwen3.6-35B-A3B CP2 allocator traces (EP1 and EP2): the head's buffers are +freed before the recompute backward allocates its workspace and adapter +gradients, and TE's first-GEMM workspaces are live at the head's peak. +""" + +import itertools +import random + +import pytest +from test_trainer_rank_adapter_gradient_memory import ( + OTHER, + POLICY, + oracle, + priced, + with_pending, + with_slot_pending, +) +from test_trainer_rank_checkpoint_memory import rank, requests + +from art.megatron.lora import LoRASlotRef +from art.trainer_rank import TrainerRank +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) +from art.trainer_rank._impl import _MemoryProfile + +MiB = 2**20 + + +def covered_rank(monkeypatch): + r = rank() + # The MoE stage covers the policy slot's recompute, as the constructor's. + monkeypatch.setattr(r, "_moe_recompute_covered_for", lambda ref: True) + return r + + +def head_oracle(chains): + """Worst adapter bytes live beyond the floor while one group's head runs. + + Any set of the other groups may have run their backward first: each holds + all its gradients and none of its boundaries. The running group holds only + its gradients outside the decoder. + """ + worst = 0 + for index, (pending, _) in enumerate(chains): + others = chains[:index] + chains[index + 1 :] + for count in range(len(others) + 1): + for done in itertools.combinations(others, count): + live = pending[-1] + sum(sum(p) - sum(b) for p, b in done) + worst = max(worst, live) + return worst + + +def test_head_adapter_gradients_match_every_order_for_random_sizes(monkeypatch): + r = rank() + generator = random.Random(2) + slots = [LoRASlotRef("checkpoint", f"slot{index}") for index in range(4)] + for _ in range(200): + layers = generator.randint(1, 5) + chains, groups, by_slot = [], [], {} + for slot in slots[: generator.randint(1, 4)]: + pending = [generator.randint(0, 30) for _ in range(layers + 1)] + boundaries = [generator.randint(0, 30) for _ in range(layers)] + by_slot[(slot,)] = pending + chains.append((pending, boundaries)) + groups.append((slot, boundaries)) + with_slot_pending(monkeypatch, r, by_slot) + assert r._checkpoint_adapter_gradient_bytes(groups, head=True) == head_oracle( + chains + ) + + +def test_one_group_head_meets_only_its_gradients_outside_the_decoder(monkeypatch): + r = rank() + with_pending(monkeypatch, r, [30] * 4 + [7]) + assert r._checkpoint_adapter_gradient_bytes(((POLICY, [2] * 4),), head=True) == 7 + # The decoder stage still meets its own gradients. + assert r._checkpoint_adapter_gradient_bytes(((POLICY, [2] * 4),)) == oracle( + [30] * 4 + [7], [2] * 4 + ) + with_slot_pending( + monkeypatch, r, {(POLICY,): [30] * 4 + [7], (OTHER,): [1] * 4 + [0]} + ) + # The policy group may run first and raise the other's head; a group whose + # boundaries outweigh its gradients never raises the policy's. + groups = ((POLICY, [2] * 4), (OTHER, [40] * 4)) + assert r._checkpoint_adapter_gradient_bytes(groups, head=True) == 127 - 8 + groups = ((POLICY, [2] * 4), (OTHER, [2] * 4)) + with_slot_pending( + monkeypatch, r, {(POLICY,): [30] * 4 + [7], (OTHER,): [0] * 4 + [0]} + ) + assert r._checkpoint_adapter_gradient_bytes(groups, head=True) == 127 - 8 + + +@pytest.mark.parametrize("head", [0, 64 * MiB, 700 * MiB, 3000 * MiB]) +def test_cost_and_estimate_price_the_larger_stage(monkeypatch, head): + r = covered_rank(monkeypatch) + n, out, signature, groups, _ = r._estimate_flat_forward(requests(67, 4096)) + values = (n, out, signature, groups, head) + retained, workspace = r._checkpoint_memory_floor(groups) + gradient = 67 * 2048 * 2 + boundary = retained // 40 + pending = [23 * MiB] * 40 + [0] + with_pending(monkeypatch, r, pending) + extra = oracle(pending, [boundary] * 40) + te = r._te_workspace_growth_bytes() + decoder = workspace + extra + stage = head + 2 * gradient + te if head else 0 + cost = priced(r, values, (POLICY, None)) + assert cost.checkpoint_input_gradient == gradient + assert cost.required == int( + (out + retained + gradient + max(decoder, stage) + COLD) * 1.1 + ) + estimate = r._estimate_required_memory_bytes_from_values( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=(POLICY, None), + head_workspace_bytes=head, + ) + assert estimate == cost.required + # Warm: no first-execution transients and no TE growth in either stage. + r._te_workspace_growth_bytes = lambda: 0 + r._memory_profiles[signature] = _MemoryProfile(bytes_per_token=1, packed_tokens=n) + decoder = r._checkpoint_memory_floor(groups)[1] + extra + assert decoder == workspace - te + extra + stage = head + 2 * gradient if head else 0 + assert priced(r, values, (POLICY, None)).required == int( + (out + retained + gradient + max(decoder, stage)) * 1.1 + ) + + +def test_head_bound_short_wave_no_longer_adds_the_adapter_extra(monkeypatch): + r = covered_rank(monkeypatch) + n, out, signature, groups, _ = r._estimate_flat_forward(requests(67, 4096)) + retained, workspace = r._checkpoint_memory_floor(groups) + head = 2 * workspace + pending = [23 * MiB] * 40 + [0] + with_pending(monkeypatch, r, pending) + extra = oracle(pending, [retained // 40] * 40) + assert head > workspace and extra > 800 * MiB + cost = priced(r, (n, out, signature, groups, head), (POLICY, None)) + unstaged = int((out + retained + head + 67 * 2048 * 2 + extra + COLD) * 1.1) + assert cost.required < unstaged + + +@pytest.mark.parametrize("uncovered", ["cp4", "dense"]) +def test_untraced_recompute_keeps_the_head_in_the_decoder_stage(monkeypatch, uncovered): + r = covered_rank(monkeypatch) + n, out, signature, groups, _ = r._estimate_flat_forward(requests(67, 4096)) + if uncovered == "cp4": + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 4, 1)) + else: + r._moe_output_bytes_per_token = r._moe_checkpoint_grad_bytes_per_token = 0 + head = 700 * MiB + retained, workspace = r._checkpoint_memory_floor(groups) + gradient = r._checkpoint_input_gradient_bytes(groups) + assert gradient == retained # One gradient per boundary. + with_pending(monkeypatch, r, ()) + cost = priced(r, (n, out, signature, groups, head), (POLICY, None)) + assert cost.checkpoint_workspace == max(workspace, head) + COLD + assert cost.required == int( + (out + retained + max(workspace, head) + gradient + COLD) * 1.1 + ) + + +def test_split_keeps_each_head_stage_beside_every_adapter_gradient(monkeypatch): + r = covered_rank(monkeypatch) + head = 700 * MiB + values = r._estimate_flat_forward(requests(67, 4096)) + n, out, signature, groups, _ = values + pending = [23 * MiB] * 40 + [0] + with_pending(monkeypatch, r, pending) + left = priced(r, (n, out, signature, groups, head), (POLICY, None)) + right = priced(r, (n, out, signature, groups, head), (POLICY, None)) + gradient = left.checkpoint_input_gradient + stage = head + 2 * gradient + r._te_workspace_growth_bytes() + assert ( + left.checkpoint_workspace + == max(r._checkpoint_memory_floor(groups)[1], stage) + COLD + ) + # A later child's decoder backward can precede an earlier child's head: + # the split charges that head stage with the shared adapter gradients. + split = TrainerRank._split_required_memory([left, right]) + assert split >= int( + ( + 2 * (left.checkpoint_retained + gradient) + + stage + + COLD + + left.checkpoint_adapter_gradient + ) + * 1.1 + ) From 06a3bb69e38ba4e716692938cbfd9773a0d25ccd Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 13:43:56 +0000 Subject: [PATCH 22/40] Price each row's RoPE and index state in the head stage The head's backward peak also holds each row's FP32 RoPE embedding and a few hundred bytes of index state (positions, row maps, CP masks, GDN exchange plans). Qwen3.6-35B-A3B CP2 multi-request traces put it at 106-181 bytes per row beyond RoPE; without it, a 4x2048 wave's head stage priced 0.6 MB below its measured peak before the safety factor. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 29 +++++++++++++++---- ...trainer_rank_checkpoint_gradient_memory.py | 10 +++++-- tests/unit/test_trainer_rank_head_memory.py | 15 +++++----- .../test_trainer_rank_head_stage_memory.py | 24 +++++++++++++-- 4 files changed, 61 insertions(+), 17 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 0b9c29789..cc1aba90a 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -130,6 +130,11 @@ class TopK: # 32 MiB transients live at its peak (Qwen3.6-35B-A3B CP2: the RoPE frequencies # and a frozen linear's output, at 2k to 20k tokens); warm waves do not. _COLD_RECOMPUTE_TRANSIENT_BYTES = 64 * 2**20 +# Per local row, index state the backward keeps beside the RoPE embedding and +# hidden-width tensors: int64 positions and row maps, CP block masks and GDN +# exchange plans. Qwen3.6-35B-A3B CP2 traces at the head's backward peak: +# 106-181 bytes per row at 3.5k-8.7k rows. +_BACKWARD_ROW_STATE_BYTES = 256 _PLANNER_REFINEMENT_BUDGET = 2_000 _LAYOUT_SELECTION_CACHE_LIMIT = 64 @@ -4776,25 +4781,39 @@ def _checkpoint_head_stage_bytes( recompute, ``gradient`` is the gradient groups' rows x 2H, and each of these is at most that: the final decoder outputs, the hidden rows each checkpointed head chunk saved (this rank's rows), and the hidden-row - gradient. Beside them are the head's own buffers, TE's cuBLAS workspaces - from the forward's first GEMMs, and the adapter gradients of groups - whose backward ran first. Qwen3.6-35B-A3B CP2 traces (EP1 and EP2) - show these terms at the head's peak. Elsewhere None: the head shares - the decoder stage. + gradient. Beside them are each row's RoPE embedding and index state, + the head's own buffers, TE's cuBLAS workspaces from the forward's first + GEMMs, and the adapter gradients of groups whose backward ran first. + Qwen3.6-35B-A3B CP2 traces (EP1 and EP2, single and multi-request + waves) show these terms at the head's peak. Elsewhere None: the head + shares the decoder stage. """ if not head_workspace_bytes or not self._checkpoint_gradient_covered( group_rows, slot_refs ): return None + rows = sum(rows for rows, grad in group_rows if grad) return ( head_workspace_bytes + 2 * gradient + + rows * self._backward_row_state_bytes() + self._te_workspace_growth_bytes() + self._checkpoint_adapter_gradient_bytes( self._checkpoint_gradient_groups(group_rows, slot_refs), head=True ) ) + def _backward_row_state_bytes(self) -> int: + """Per local row, what the backward keeps beside hidden-width tensors. + + The FP32 RoPE embedding at the rotary width (256 bytes for + Qwen3.6-35B-A3B) and ``_BACKWARD_ROW_STATE_BYTES`` of index state. + """ + rope = getattr(_language_model(self.runtime.model[0]), "rotary_pos_emb", None) + frequencies = getattr(rope, "inv_freq", None) + width = 2 * frequencies.numel() if isinstance(frequencies, torch.Tensor) else 0 + return width * 4 + _BACKWARD_ROW_STATE_BYTES + def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( packed_tokens=plan.packed_tokens, diff --git a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py index d5bf9e3a3..c27ca5e01 100644 --- a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py @@ -141,8 +141,14 @@ def test_gradient_is_not_absorbed_by_larger_head_workspace(): cost = price(r, (n, out, sig, groups, head)) boundaries = 67 * 40 * 2048 * 2 gradient = 67 * 2048 * 2 - # The head stage: two more gradient-row terms and TE's first-GEMM workspaces. - stage = head + 2 * gradient + r._te_workspace_growth_bytes() + # The head stage: two more gradient-row terms, each row's RoPE and index + # state, and TE's first-GEMM workspaces. + stage = ( + head + + 2 * gradient + + 67 * r._backward_row_state_bytes() + + r._te_workspace_growth_bytes() + ) assert cost.checkpoint_workspace == stage + COLD assert cost.required == int((out + stage + COLD + boundaries + gradient) * 1.1) assert cost.retained == int((out + head + boundaries) * 1.1) diff --git a/tests/unit/test_trainer_rank_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py index 26f44aa7e..f3f24928e 100644 --- a/tests/unit/test_trainer_rank_head_memory.py +++ b/tests/unit/test_trainer_rank_head_memory.py @@ -151,11 +151,13 @@ def test_outputs_retention_and_empirical_peak_are_counted_once(): cost = r._plan_cost(plan) assert cost.retained == int((plan.output_bytes + retained + head) * 1.1) # Unprofiled head stage: the head workspace, the final output, saved - # selected rows and hidden-row gradient beside the incoming one, TE's - # first-GEMM workspaces and the first execution's transients. + # selected rows and hidden-row gradient beside the incoming one, each + # row's RoPE and index state, TE's first-GEMM workspaces and the first + # execution's transients. + state = 512 * r._backward_row_state_bytes() te = r._te_workspace_growth_bytes() assert cost.required == int( - (plan.output_bytes + retained + 3 * gradient + head + te + COLD) * 1.1 + (plan.output_bytes + retained + 3 * gradient + state + head + te + COLD) * 1.1 ) r._memory_profiles[plan.signature] = _MemoryProfile( bytes_per_token=2_000_000, @@ -322,11 +324,10 @@ def test_target_backward_refuses_budget_below_logits_and_both_gradients(rows): plan = r._plan_flat_forward([request(rows, grad=True)]) retained, _ = r._checkpoint_memory_floor(r._plan_group_rows(plan)) gradient = rows * 2048 * 2 + stage = 3 * gradient + rows * r._backward_row_state_bytes() dense = min(rows, 512) * 248320 * 2 - before = int((plan.output_bytes + retained + 3 * gradient + 2 * dense + COLD) * 1.1) - expected = int( - (plan.output_bytes + retained + 3 * gradient + 3 * dense + COLD) * 1.1 - ) + before = int((plan.output_bytes + retained + stage + 2 * dense + COLD) * 1.1) + expected = int((plan.output_bytes + retained + stage + 3 * dense + COLD) * 1.1) r._available_memory_bytes = lambda: (before + expected) // 2 check = r._memory_check(plan) assert check.estimated_required_bytes == expected diff --git a/tests/unit/test_trainer_rank_head_stage_memory.py b/tests/unit/test_trainer_rank_head_stage_memory.py index 2d67b68c5..7b75e2fcf 100644 --- a/tests/unit/test_trainer_rank_head_stage_memory.py +++ b/tests/unit/test_trainer_rank_head_stage_memory.py @@ -106,8 +106,9 @@ def test_cost_and_estimate_price_the_larger_stage(monkeypatch, head): with_pending(monkeypatch, r, pending) extra = oracle(pending, [boundary] * 40) te = r._te_workspace_growth_bytes() + state = 67 * r._backward_row_state_bytes() decoder = workspace + extra - stage = head + 2 * gradient + te if head else 0 + stage = head + 2 * gradient + state + te if head else 0 cost = priced(r, values, (POLICY, None)) assert cost.checkpoint_input_gradient == gradient assert cost.required == int( @@ -128,7 +129,7 @@ def test_cost_and_estimate_price_the_larger_stage(monkeypatch, head): r._memory_profiles[signature] = _MemoryProfile(bytes_per_token=1, packed_tokens=n) decoder = r._checkpoint_memory_floor(groups)[1] + extra assert decoder == workspace - te + extra - stage = head + 2 * gradient if head else 0 + stage = head + 2 * gradient + state if head else 0 assert priced(r, values, (POLICY, None)).required == int( (out + retained + gradient + max(decoder, stage)) * 1.1 ) @@ -178,7 +179,12 @@ def test_split_keeps_each_head_stage_beside_every_adapter_gradient(monkeypatch): left = priced(r, (n, out, signature, groups, head), (POLICY, None)) right = priced(r, (n, out, signature, groups, head), (POLICY, None)) gradient = left.checkpoint_input_gradient - stage = head + 2 * gradient + r._te_workspace_growth_bytes() + stage = ( + head + + 2 * gradient + + 67 * r._backward_row_state_bytes() + + r._te_workspace_growth_bytes() + ) assert ( left.checkpoint_workspace == max(r._checkpoint_memory_floor(groups)[1], stage) + COLD @@ -195,3 +201,15 @@ def test_split_keeps_each_head_stage_beside_every_adapter_gradient(monkeypatch): ) * 1.1 ) + + +def test_each_row_keeps_its_rope_embedding_and_index_state(): + import torch + + r = rank() + model = r.runtime.model[0] + assert r._backward_row_state_bytes() == 256 + # Qwen3.6-35B-A3B: a 64-wide rotary embedding (32 frequencies), FP32. + model.rotary_pos_emb = torch.nn.Module() + model.rotary_pos_emb.inv_freq = torch.ones(32) + assert r._backward_row_state_bytes() == 64 * 4 + 256 From e1982396d57cd0dbb53b5744cfc151e96add426a Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 14:25:39 +0000 Subject: [PATCH 23/40] Stage only traced head backwards; keep the unstaged price for others The head stage relies on the head workspace bounding the head's backward buffers. Qwen3.6-35B-A3B CP2 traces establish that only for target-only requests on the standard logit scale through the fused Triton statistics. A top-k-only head, for one, keeps its recomputed logits beside their gradient, twice the one buffer its workspace prices, which the adapter extra used to absorb. Any other head, and CP sizes other than two, keep the unstaged price as a floor beneath the staged one. Bounds that leave the projected rows open price both ways: the lower to reject, the higher to accept. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 226 ++++++++++++++---- .../test_trainer_rank_head_stage_memory.py | 123 ++++++++-- 2 files changed, 285 insertions(+), 64 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index cc1aba90a..4287b580c 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -3729,6 +3729,7 @@ def _split_chunk_lower_cost( packed_tokens = 0 unshared_packed_tokens = 0 head_workspace_bytes = 0 + head_traced: list[bool | None] = [] group_rows: list[tuple[int, bool]] = [] for (_slot, grad_enabled), group_indices in groups: estimated = estimate_prefix_tree_packed_tokens( @@ -3742,15 +3743,24 @@ def _split_chunk_lower_cost( cp = max(1, self._topology_key()[2]) group_rows.append((-(-physical_rows // cp), grad_enabled)) head_requests = tuple(requests[index] for index in group_indices) + lower = self._head_projection_rows(head_requests, lower_bound=True) head_workspace_bytes = max( head_workspace_bytes, self._group_head_workspace_bytes( - self._head_projection_rows(head_requests, lower_bound=True), + lower, head_requests, grad_enabled=grad_enabled, lower_bound=True, ), ) + if grad_enabled: + head_traced.append( + self._head_backward_traced( + head_requests, + lower, + self._head_projection_rows(head_requests), + ) + ) unshared_packed_tokens += self._physical_tokens( sum(int(rows[index].numel()) for index in group_indices) ) @@ -3762,17 +3772,27 @@ def _split_chunk_lower_cost( slot_groups=tuple(key for key, _ in groups), ) logical_tokens = _active_logical_tokens(requests) - cost = self._subforward_cost( - packed_tokens=packed_tokens, - output_bytes=output_bytes, - signature=signature, - logical_tokens=logical_tokens, - group_rows=tuple(group_rows), - slot_refs=tuple(ref for (ref, _), _ in groups), - head_workspace_bytes=head_workspace_bytes, - # The average CP load is an optimistic bound, not an admission cost. - retained_tokens=(packed_tokens + signature.topology[2] - 1) - // signature.topology[2], + # Whether the exact plan stages its head can depend on projected rows + # these bounds leave open; the lower of both prices bounds either. + cost = min( + ( + self._subforward_cost( + packed_tokens=packed_tokens, + output_bytes=output_bytes, + signature=signature, + logical_tokens=logical_tokens, + group_rows=tuple(group_rows), + slot_refs=tuple(ref for (ref, _), _ in groups), + head_workspace_bytes=head_workspace_bytes, + head_backward_traced=traced, + # The average CP load is an optimistic bound, not an + # admission cost. + retained_tokens=(packed_tokens + signature.topology[2] - 1) + // signature.topology[2], + ) + for traced in _traced_states(head_traced) + ), + key=lambda cost: cost.required, ) profile = self._memory_profiles.get(signature) if ( @@ -4009,18 +4029,7 @@ def _group_head_workspace_bytes( or not any(request.target_tokens is not None for request in requests) ): return dense - from megatron.core.models.common.language_module.language_module import ( - LanguageModule, - ) - - model = _language_model(self.runtime.model[0]) - scale = getattr(model, "_scale_logits", None) - if ( - type(scale) is MethodType - and scale.__self__ is model - and scale.__func__ is LanguageModule._scale_logits - and getattr(model.config, "use_mup", None) is False - ): + if self._standard_logit_scale(): # IndexBackward's dense result overlaps saved logits and grad_logits. # The FP32 fallback already exceeds this three-buffer component. target_dense = ( @@ -4037,6 +4046,79 @@ def _group_head_workspace_bytes( return max(dense, 3 * target_dense) return dense + def _standard_logit_scale(self) -> bool: + """Whether the head scales logits with Megatron's own, un-muP'd method.""" + from megatron.core.models.common.language_module.language_module import ( + LanguageModule, + ) + + model = _language_model(self.runtime.model[0]) + scale = getattr(model, "_scale_logits", None) + return ( + type(scale) is MethodType + and scale.__self__ is model + and scale.__func__ is LanguageModule._scale_logits + and getattr(model.config, "use_mup", None) is False + ) + + def _head_backward_traced( + self, + requests: Sequence[AnyForwardInput], + rows: int, + upper_rows: int | None = None, + ) -> bool | None: + """Whether a gradient group's head backward is the traced one. + + The head stage (``_checkpoint_head_stage_bytes``) relies on + ``_group_head_workspace_bytes`` bounding the head's backward buffers. + Qwen3.6-35B-A3B CP2 traces establish that for target-only requests on + the standard logit scale through the fused Triton statistics, which + need at least ``ART_TRAINER_RANK_TRITON_MIN_ROWS`` rows in the first + projected chunk. Top-k, logits and hidden-state outputs keep further + dense gradients, and the FP32 fallback wider copies; other CP sizes + are untraced. ``rows`` bounds the first chunk's rows from below, + ``upper_rows`` from above: None when they straddle the threshold. + """ + if ( + not requests + or any( + request.target_tokens is None + or request.top_k is not None + or request.logits + or request.hidden_states + for request in requests + ) + or os.environ.get("ART_TRAINER_RANK_TRITON_TOPK", "1").lower() + in {"0", "false"} + or self._topology_key()[2] != 2 + or not self._head_workspace_bytes(1) + or not self._standard_logit_scale() + ): + return False + minimum = int(os.environ.get("ART_TRAINER_RANK_TRITON_MIN_ROWS", "64")) + if rows >= minimum: + return True + if upper_rows is None or upper_rows < minimum: + return False + return None + + def _plan_head_backward_traced(self, plan: _FlatForwardPlan) -> bool: + """Every gradient group's head backward is traced (``_head_backward_traced``).""" + traced = [ + self._head_backward_traced( + requests, + self._head_projection_rows( + requests, + positions=group.packed.positions_by_sequence, + lower_bound=True, + ), + ) + for group in plan.groups + if group.grad_enabled + for requests in (tuple(item.request for item in group.items),) + ] + return bool(traced) and all(state is True for state in traced) + def _plan_head_workspace_bytes(self, plan: _FlatForwardPlan) -> int: peak = 0 for group in plan.groups: @@ -4786,7 +4868,9 @@ def _checkpoint_head_stage_bytes( GEMMs, and the adapter gradients of groups whose backward ran first. Qwen3.6-35B-A3B CP2 traces (EP1 and EP2, single and multi-request waves) show these terms at the head's peak. Elsewhere None: the head - shares the decoder stage. + shares the decoder stage. That bounds heads whose backward is the + traced one (``_head_backward_traced``); callers keep the unstaged price + beside any other, whose buffers can exceed ``head_workspace_bytes``. """ if not head_workspace_bytes or not self._checkpoint_gradient_covered( group_rows, slot_refs @@ -4825,6 +4909,7 @@ def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: group_routed_rows=self._plan_group_routed_rows(plan), slot_refs=tuple(g.slot_ref for g in plan.groups), head_workspace_bytes=self._plan_head_workspace_bytes(plan), + head_backward_traced=self._plan_head_backward_traced(plan), checkpoint_floor=_gdn_memory.plan_floor(self, plan), retained_tokens=self._plan_retained_tokens(plan), hybridep_growth_bytes=self._plan_hybridep_growth_bytes(plan), @@ -4842,6 +4927,7 @@ def _subforward_cost( group_routed_rows: tuple[int, ...] | None = None, slot_refs: tuple["LoRASlotRef | None", ...] | None = None, head_workspace_bytes: int = 0, + head_backward_traced: bool = False, checkpoint_floor: tuple[int, int] = (0, 0), retained_tokens: int | None = None, hybridep_growth_bytes: int = 0, @@ -4898,14 +4984,16 @@ def _subforward_cost( head_stage = self._checkpoint_head_stage_bytes( head_workspace_bytes, gradient, group_rows, slot_refs ) - if head_stage is None: - peak = checkpoint_workspace + adapter_gradient - else: + peak = checkpoint_workspace + adapter_gradient + if head_stage is not None: # A split adds its children's adapter gradients to the largest # workspace, and one child's head can follow another's decoder # backward: keep the whole head stage there. checkpoint_workspace = max(decoder_workspace, head_stage) - peak = max(decoder_workspace + adapter_gradient, head_stage) + # An untraced head's buffers can exceed head_workspace_bytes: + # it keeps the unstaged price as a floor. + staged = max(decoder_workspace + adapter_gradient, head_stage) + peak = staged if head_backward_traced else max(peak, staged) if self._memory_profiles.get(signature) is None: checkpoint_workspace += _COLD_RECOMPUTE_TRANSIENT_BYTES peak += _COLD_RECOMPUTE_TRANSIENT_BYTES @@ -5925,11 +6013,13 @@ def estimate(width: int) -> tuple[_MemoryCheck, bool, bool] | None: indices, local_inputs = local_slice(width) local_requests = list(_flatten(local_inputs)) cheap_segments: list[int] = [] + cheap_traced: list[bool | None] = [] values = self._estimate_flat_forward( local_requests, checkpoint=checkpoint, sync_planning_errors=True, gdn_segments=cheap_segments, + head_traced=cheap_traced, ) if not self._all_ranks_true(values is not None): estimates[width] = None @@ -5945,18 +6035,27 @@ def priced( head_workspace_bytes: int, *, gdn_segments: int, + head_traced: Sequence[bool | None], + lower: bool, ) -> tuple[_MemoryCheck, int, int, _MemorySignature]: with self._planning_status(True): - required = self._estimate_required_memory_bytes_from_values( - packed_tokens=packed_tokens, - output_bytes=output_bytes, - signature=signature, - logical_tokens=logical_tokens, - # Gradient groups' segments: exact layouts' counts, else - # a bound matching the estimate's (_estimate_flat_forward). - gdn_segments=gdn_segments, - group_rows=group_rows, - head_workspace_bytes=head_workspace_bytes, + # A bound over layouts whose head may or may not stage: + # the higher price to accept, the lower to reject. + required = (min if lower else max)( + self._estimate_required_memory_bytes_from_values( + packed_tokens=packed_tokens, + output_bytes=output_bytes, + signature=signature, + logical_tokens=logical_tokens, + # Gradient groups' segments: exact layouts' counts, + # else a bound matching the estimate's + # (_estimate_flat_forward). + gdn_segments=gdn_segments, + group_rows=group_rows, + head_workspace_bytes=head_workspace_bytes, + head_backward_traced=traced, + ) + for traced in _traced_states(head_traced) ) return ( self._memory_check_required(required, sync_across_dp=True), @@ -5969,6 +6068,7 @@ def priced_estimate( *, exact: bool, memory_minimal: bool ) -> tuple[_MemoryCheck, int, int, _MemorySignature] | None: segments: list[int] = [] + traced: list[bool | None] = [] estimated = self._estimate_flat_forward( local_requests, checkpoint=checkpoint, @@ -5976,11 +6076,18 @@ def priced_estimate( memory_minimal=memory_minimal, sync_planning_errors=True, gdn_segments=segments, + head_traced=traced, ) return ( None if estimated is None - else priced(*estimated, gdn_segments=sum(segments)) + else priced( + *estimated, + gdn_segments=sum(segments), + head_traced=traced, + # Only the cheap full-sharing count rejects. + lower=memory_minimal and not exact, + ) ) def trusted(packed_tokens: int, signature: _MemorySignature) -> bool: @@ -5994,7 +6101,12 @@ def trusted(packed_tokens: int, signature: _MemorySignature) -> bool: # reject on memory, or when it would reject on profile trust while # a profile exists — the selected layout may be far smaller than # the bound and squarely inside the profiled regime. - selected = priced(*values, gdn_segments=sum(cheap_segments)) + selected = priced( + *values, + gdn_segments=sum(cheap_segments), + head_traced=cheap_traced, + lower=False, + ) profiled = self._all_ranks_true(selected[3] in self._memory_profiles) needs_exact = not selected[0].fits or ( profiled and not trusted(selected[1], selected[3]) @@ -6668,6 +6780,7 @@ def _estimate_flat_forward( memory_minimal: bool = False, sync_planning_errors: bool = False, gdn_segments: list[int] | None = None, + head_traced: list[bool | None] | None = None, ) -> tuple[int, int, _MemorySignature, tuple[tuple[int, bool], ...], int] | None: """Estimate packed tokens for width probing. @@ -6683,6 +6796,8 @@ def _estimate_flat_forward( ``gdn_segments`` receives each gradient group's segment count: exact layouts' actual counts; in cheap mode, the same kind of bound as the token count (twice the requests, as a radix tree has fewer, or one). + ``head_traced`` receives each gradient group's ``_head_backward_traced`` + over its projected-row bounds. """ if sync_planning_errors: @@ -6731,6 +6846,10 @@ def _estimate_flat_forward( head_requests = tuple(requests[index] for index in group_indices) lower = self._head_projection_rows(head_requests, lower_bound=True) upper = self._head_projection_rows(head_requests) + if grad_enabled and head_traced is not None: + head_traced.append( + self._head_backward_traced(head_requests, lower, upper) + ) if exact: tree, layout = self._select_group_layout( tuple( @@ -7768,6 +7887,7 @@ def _memory_check( group_routed_rows=self._plan_group_routed_rows(forward), slot_refs=tuple(g.slot_ref for g in forward.groups), head_workspace_bytes=self._plan_head_workspace_bytes(forward), + head_backward_traced=self._plan_head_backward_traced(forward), checkpoint_floor=_gdn_memory.plan_floor(self, forward), retained_tokens=self._plan_retained_tokens(forward), ) + int(self._plan_hybridep_growth_bytes(forward) * _MEMORY_SAFETY_FACTOR) @@ -8615,6 +8735,7 @@ def _estimate_required_memory_bytes_from_values( group_routed_rows: tuple[int, ...] | None = None, slot_refs: tuple["LoRASlotRef | None", ...] | None = None, head_workspace_bytes: int = 0, + head_backward_traced: bool = False, checkpoint_floor: tuple[int, int] = (0, 0), retained_tokens: int | None = None, include_checkpoint_input_gradient: bool = True, @@ -8722,11 +8843,12 @@ def _estimate_required_memory_bytes_from_values( head_stage = self._checkpoint_head_stage_bytes( head_workspace_bytes, gradient, group_rows, slot_refs ) - peak = gradient + ( - peak + adapter_gradient - if head_stage is None - else max(decoder_workspace + adapter_gradient, head_stage) - ) + peak += adapter_gradient + if head_stage is not None: + # An untraced head keeps the unstaged price as a floor. + staged = max(decoder_workspace + adapter_gradient, head_stage) + peak = staged if head_backward_traced else max(peak, staged) + peak += gradient if profiled is None and any(grad for _, grad in group_rows): peak += _COLD_RECOMPUTE_TRANSIENT_BYTES static_compute = max(static_compute, max(retained, checkpoint_floor[0]) + peak) @@ -9904,6 +10026,18 @@ def _active_logical_tokens(requests: Sequence[AnyForwardInput]) -> int: _PACKED_PRICED_MIN_REQUEST_TOKENS = 64 +def _traced_states(traced: Sequence[bool | None]) -> tuple[bool, ...]: + """Head stagings a plan's gradient groups allow (``_head_backward_traced``). + + Staged only when every group is traced; both when any is undecided. + """ + if not traced or any(state is False for state in traced): + return (False,) + if all(state is True for state in traced): + return (True,) + return (True, False) + + def _packed_priced(signature: "_MemorySignature", one_layer_recompute: bool) -> bool: # Measured only under one-layer full recompute: other recompute modes keep # more GDN states and activations live per segment and logical row. diff --git a/tests/unit/test_trainer_rank_head_stage_memory.py b/tests/unit/test_trainer_rank_head_stage_memory.py index 7b75e2fcf..742f5a28a 100644 --- a/tests/unit/test_trainer_rank_head_stage_memory.py +++ b/tests/unit/test_trainer_rank_head_stage_memory.py @@ -5,6 +5,7 @@ gradients, and TE's first-GEMM workspaces are live at the head's peak. """ +from dataclasses import replace import itertools import random @@ -94,6 +95,21 @@ def test_one_group_head_meets_only_its_gradients_outside_the_decoder(monkeypatch assert r._checkpoint_adapter_gradient_bytes(groups, head=True) == 127 - 8 +def staged(r, values, slot_refs): + """``priced`` for a head whose backward is the traced one.""" + n, out, signature, groups, head = values + return r._subforward_cost( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=slot_refs, + head_workspace_bytes=head, + head_backward_traced=True, + ) + + @pytest.mark.parametrize("head", [0, 64 * MiB, 700 * MiB, 3000 * MiB]) def test_cost_and_estimate_price_the_larger_stage(monkeypatch, head): r = covered_rank(monkeypatch) @@ -109,7 +125,7 @@ def test_cost_and_estimate_price_the_larger_stage(monkeypatch, head): state = 67 * r._backward_row_state_bytes() decoder = workspace + extra stage = head + 2 * gradient + state + te if head else 0 - cost = priced(r, values, (POLICY, None)) + cost = staged(r, values, (POLICY, None)) assert cost.checkpoint_input_gradient == gradient assert cost.required == int( (out + retained + gradient + max(decoder, stage) + COLD) * 1.1 @@ -122,15 +138,23 @@ def test_cost_and_estimate_price_the_larger_stage(monkeypatch, head): group_rows=groups, slot_refs=(POLICY, None), head_workspace_bytes=head, + head_backward_traced=True, ) assert estimate == cost.required + # An untraced head keeps the unstaged price as a floor. + unstaged = max(workspace, head) + extra + untraced = priced(r, values, (POLICY, None)) + assert untraced.required == int( + (out + retained + gradient + max(unstaged, decoder, stage) + COLD) * 1.1 + ) + assert untraced.checkpoint_workspace == cost.checkpoint_workspace # Warm: no first-execution transients and no TE growth in either stage. r._te_workspace_growth_bytes = lambda: 0 r._memory_profiles[signature] = _MemoryProfile(bytes_per_token=1, packed_tokens=n) decoder = r._checkpoint_memory_floor(groups)[1] + extra assert decoder == workspace - te + extra stage = head + 2 * gradient + state if head else 0 - assert priced(r, values, (POLICY, None)).required == int( + assert staged(r, values, (POLICY, None)).required == int( (out + retained + gradient + max(decoder, stage)) * 1.1 ) @@ -144,12 +168,15 @@ def test_head_bound_short_wave_no_longer_adds_the_adapter_extra(monkeypatch): with_pending(monkeypatch, r, pending) extra = oracle(pending, [retained // 40] * 40) assert head > workspace and extra > 800 * MiB - cost = priced(r, (n, out, signature, groups, head), (POLICY, None)) + values = (n, out, signature, groups, head) unstaged = int((out + retained + head + 67 * 2048 * 2 + extra + COLD) * 1.1) - assert cost.required < unstaged + assert staged(r, values, (POLICY, None)).required < unstaged + # Untraced (e.g. a top-k-only head, whose backward keeps its recomputed + # logits beside their gradient): never below the unstaged price. + assert priced(r, values, (POLICY, None)).required == unstaged -@pytest.mark.parametrize("uncovered", ["cp4", "dense"]) +@pytest.mark.parametrize("uncovered", ["cp4", "uncovered_moe"]) def test_untraced_recompute_keeps_the_head_in_the_decoder_stage(monkeypatch, uncovered): r = covered_rank(monkeypatch) n, out, signature, groups, _ = r._estimate_flat_forward(requests(67, 4096)) @@ -162,7 +189,8 @@ def test_untraced_recompute_keeps_the_head_in_the_decoder_stage(monkeypatch, unc gradient = r._checkpoint_input_gradient_bytes(groups) assert gradient == retained # One gradient per boundary. with_pending(monkeypatch, r, ()) - cost = priced(r, (n, out, signature, groups, head), (POLICY, None)) + # Even for a traced head. + cost = staged(r, (n, out, signature, groups, head), (POLICY, None)) assert cost.checkpoint_workspace == max(workspace, head) + COLD assert cost.required == int( (out + retained + max(workspace, head) + gradient + COLD) * 1.1 @@ -176,8 +204,8 @@ def test_split_keeps_each_head_stage_beside_every_adapter_gradient(monkeypatch): n, out, signature, groups, _ = values pending = [23 * MiB] * 40 + [0] with_pending(monkeypatch, r, pending) - left = priced(r, (n, out, signature, groups, head), (POLICY, None)) - right = priced(r, (n, out, signature, groups, head), (POLICY, None)) + left = staged(r, (n, out, signature, groups, head), (POLICY, None)) + right = staged(r, (n, out, signature, groups, head), (OTHER, None)) gradient = left.checkpoint_input_gradient stage = ( head @@ -189,18 +217,61 @@ def test_split_keeps_each_head_stage_beside_every_adapter_gradient(monkeypatch): left.checkpoint_workspace == max(r._checkpoint_memory_floor(groups)[1], stage) + COLD ) - # A later child's decoder backward can precede an earlier child's head: - # the split charges that head stage with the shared adapter gradients. + # The left child's head can run after the right child's decoder backward: + # both children's boundaries and incoming gradients, the left head stage + # and every adapter gradient the right one allocated (distinct slots). + after_right = ( + left.checkpoint_retained + + right.checkpoint_retained + + 2 * gradient + + stage + + COLD + + right.checkpoint_adapter_gradient + ) split = TrainerRank._split_required_memory([left, right]) + assert split >= int(after_right * 1.1) assert split >= int( - ( - 2 * (left.checkpoint_retained + gradient) - + stage - + COLD - + left.checkpoint_adapter_gradient - ) - * 1.1 - ) + (after_right + left.checkpoint_adapter_gradient) * 1.1 + ) # Either may run first, and distinct slots each allocate their own. + + +def test_traced_head_backward_is_target_only_on_the_traced_path(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + r = head_rank() + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + target = request(512, grad=True) + assert r._head_backward_traced([target], 512) is True + # Top-k, logits and hidden-state outputs keep further dense gradients. + for extra in ({"top_k": 2}, {"logits": True}, {"hidden_states": True}): + other = replace(target, **extra) + assert r._head_backward_traced([target, other], 512) is False + topk_only = replace(target, target_tokens=None, top_k=2) + assert r._head_backward_traced([topk_only], 512) is False + # The fused statistics need 64 rows; bounds that straddle them are open. + assert r._head_backward_traced([target], 63) is False + assert r._head_backward_traced([target], 63, 64) is None + assert r._head_backward_traced([target], 32, 63) is False + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_TOPK", "0") + assert r._head_backward_traced([target], 512) is False + monkeypatch.delenv("ART_TRAINER_RANK_TRITON_TOPK") + # Only CP2 was traced. + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 1, 1)) + assert r._head_backward_traced([target], 512) is False + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + r.runtime.model[0].config.use_mup = True + assert r._head_backward_traced([target], 512) is False + + +def test_undecided_heads_bound_both_ways(): + from art.trainer_rank._impl import _traced_states + + assert _traced_states([]) == (False,) + assert _traced_states([True, True]) == (True,) + assert _traced_states([True, False, None]) == (False,) + # Lower bounds take the cheaper, acceptance the dearer of both. + assert _traced_states([True, None]) == (True, False) def test_each_row_keeps_its_rope_embedding_and_index_state(): @@ -213,3 +284,19 @@ def test_each_row_keeps_its_rope_embedding_and_index_state(): model.rotary_pos_emb = torch.nn.Module() model.rotary_pos_emb.inv_freq = torch.ones(32) assert r._backward_row_state_bytes() == 64 * 4 + 256 + + +def test_plan_stages_only_traced_gradient_heads(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + r = head_rank() + target = request(512, grad=True) + topk_only = replace(target, target_tokens=None, top_k=2) + plans = [r._plan_flat_forward([request]) for request in (target, topk_only)] + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + # A top-k-only head's backward keeps its recomputed logits beside their + # gradient, twice the one buffer its workspace prices: never staged. + assert [r._plan_head_backward_traced(plan) for plan in plans] == [True, False] + no_grad = r._plan_flat_forward([request(512)]) + assert r._plan_head_backward_traced(no_grad) is False From 6a249ec434b5d24cd915677a1f16f366ced4ba71 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 14:45:34 +0000 Subject: [PATCH 24/40] Trust staging only after the fused head statistics have run A failed fused kernel falls back to FP32 statistics silently, and a CP rank can project fewer rows than the fused minimum however many the group has, so the planner cannot know the head's path from the plan alone. Stage a head only once the fused statistics have run in this process with no error fallback, and cover a chunk below the minimum with the FP32 fallback's buffers (at most nine BF16 logits-sized, from the code). Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 33 +++++++++++++++-- tests/unit/test_trainer_rank_head_memory.py | 8 +++- .../test_trainer_rank_head_stage_memory.py | 37 +++++++++++++++++++ 3 files changed, 73 insertions(+), 5 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 4287b580c..672e924b0 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -135,6 +135,16 @@ class TopK: # exchange plans. Qwen3.6-35B-A3B CP2 traces at the head's backward peak: # 106-181 bytes per row at 3.5k-8.7k rows. _BACKWARD_ROW_STATE_BYTES = 256 +# A head chunk below the fused statistics' row minimum takes the FP32 fallback +# (_vocab_parallel_log_z): the BF16 logits, their FP32 copy, the shifted copy +# and its saved exponent in the recompute, then the exponent's gradient, its +# BF16 cast and the target gather's gradient in backward. At most this many +# BF16 logits-sized buffers of that chunk (about 18 bytes per logit); derived +# from the code, not traced. +_HEAD_FALLBACK_BUFFERS = 9 +# Whether the fused head statistics have run in this process, and whether any +# call fell back to FP32 after an error: staging trusts only a proven path. +_TRITON_STATS_STATE = {"succeeded": False, "failed": False} _PLANNER_REFINEMENT_BUDGET = 2_000 _LAYOUT_SELECTION_CACHE_LIMIT = 64 @@ -4076,8 +4086,12 @@ def _head_backward_traced( need at least ``ART_TRAINER_RANK_TRITON_MIN_ROWS`` rows in the first projected chunk. Top-k, logits and hidden-state outputs keep further dense gradients, and the FP32 fallback wider copies; other CP sizes - are untraced. ``rows`` bounds the first chunk's rows from below, + are untraced. The fused statistics must have run in this process and + never fallen back after an error: a failed kernel takes the FP32 path + silently. ``rows`` bounds the first chunk's rows from below, ``upper_rows`` from above: None when they straddle the threshold. + A CP rank projecting fewer rows falls back on its own; the head stage + prices that (``_HEAD_FALLBACK_BUFFERS``). """ if ( not requests @@ -4090,6 +4104,8 @@ def _head_backward_traced( ) or os.environ.get("ART_TRAINER_RANK_TRITON_TOPK", "1").lower() in {"0", "false"} + or not _TRITON_STATS_STATE["succeeded"] + or _TRITON_STATS_STATE["failed"] or self._topology_key()[2] != 2 or not self._head_workspace_bytes(1) or not self._standard_logit_scale() @@ -4877,8 +4893,16 @@ def _checkpoint_head_stage_bytes( ): return None rows = sum(rows for rows, grad in group_rows if grad) + # A rank's chunk below the fused minimum (CP splits the projected rows + # unevenly) takes the FP32 fallback. + fallback = _HEAD_FALLBACK_BUFFERS * self._head_workspace_bytes( + min( + int(os.environ.get("ART_TRAINER_RANK_TRITON_MIN_ROWS", "64")) - 1, + _HEAD_CHUNK_TOKENS, + ) + ) return ( - head_workspace_bytes + max(head_workspace_bytes, fallback) + 2 * gradient + rows * self._backward_row_state_bytes() + self._te_workspace_growth_bytes() @@ -10604,11 +10628,14 @@ def _try_triton_stats( try: from art.trainer_rank import topk - return getattr(topk, name)(local_logits, **kwargs) + result = getattr(topk, name)(local_logits, **kwargs) except Exception: + _TRITON_STATS_STATE["failed"] = True if os.environ.get("ART_TRAINER_RANK_TRITON_TOPK", "1").lower() == "strict": raise return None + _TRITON_STATS_STATE["succeeded"] = True + return result def _vocab_parallel_topk_from_local( diff --git a/tests/unit/test_trainer_rank_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py index f3f24928e..569037e40 100644 --- a/tests/unit/test_trainer_rank_head_memory.py +++ b/tests/unit/test_trainer_rank_head_memory.py @@ -317,10 +317,14 @@ def test_tied_standard_head_weight_uses_the_same_capacity(): @pytest.mark.parametrize("rows", [128, 512]) -def test_target_backward_refuses_budget_below_logits_and_both_gradients(rows): +def test_target_backward_refuses_budget_below_logits_and_both_gradients( + monkeypatch, rows +): r = rank() - # Isolate the head term from TE's one-time cuBLAS workspace growth. + # Isolate the head term from TE's one-time cuBLAS workspace growth, and + # the three target-backward buffers from the small-chunk fallback cover. r._te_workspace_growth_bytes = lambda: 0 + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", "1") plan = r._plan_flat_forward([request(rows, grad=True)]) retained, _ = r._checkpoint_memory_floor(r._plan_group_rows(plan)) gradient = rows * 2048 * 2 diff --git a/tests/unit/test_trainer_rank_head_stage_memory.py b/tests/unit/test_trainer_rank_head_stage_memory.py index 742f5a28a..4c76f6e86 100644 --- a/tests/unit/test_trainer_rank_head_stage_memory.py +++ b/tests/unit/test_trainer_rank_head_stage_memory.py @@ -239,9 +239,16 @@ def test_traced_head_backward_is_target_only_on_the_traced_path(monkeypatch): from test_trainer_rank_head_memory import rank as head_rank from test_trainer_rank_head_memory import request + from art.trainer_rank import _impl + r = head_rank() monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) target = request(512, grad=True) + # Until the fused statistics have run in this process, nothing is traced. + state = {"succeeded": False, "failed": False} + monkeypatch.setattr(_impl, "_TRITON_STATS_STATE", state) + assert r._head_backward_traced([target], 512) is False + state["succeeded"] = True assert r._head_backward_traced([target], 512) is True # Top-k, logits and hidden-state outputs keep further dense gradients. for extra in ({"top_k": 2}, {"logits": True}, {"hidden_states": True}): @@ -260,6 +267,10 @@ def test_traced_head_backward_is_target_only_on_the_traced_path(monkeypatch): monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 1, 1)) assert r._head_backward_traced([target], 512) is False monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + # One error sent a kernel to the FP32 fallback silently: never again. + state["failed"] = True + assert r._head_backward_traced([target], 512) is False + state["failed"] = False r.runtime.model[0].config.use_mup = True assert r._head_backward_traced([target], 512) is False @@ -290,6 +301,11 @@ def test_plan_stages_only_traced_gradient_heads(monkeypatch): from test_trainer_rank_head_memory import rank as head_rank from test_trainer_rank_head_memory import request + from art.trainer_rank import _impl + + monkeypatch.setattr( + _impl, "_TRITON_STATS_STATE", {"succeeded": True, "failed": False} + ) r = head_rank() target = request(512, grad=True) topk_only = replace(target, target_tokens=None, top_k=2) @@ -300,3 +316,24 @@ def test_plan_stages_only_traced_gradient_heads(monkeypatch): assert [r._plan_head_backward_traced(plan) for plan in plans] == [True, False] no_grad = r._plan_flat_forward([request(512)]) assert r._plan_head_backward_traced(no_grad) is False + + +def test_head_stage_covers_a_small_chunks_fp32_fallback(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + + r = head_rank() + monkeypatch.setattr(r, "_moe_recompute_covered_for", lambda ref: True) + monkeypatch.setattr(r, "_checkpoint_gradient_covered", lambda *a, **k: True) + r._te_workspace_growth_bytes = lambda: 0 + groups = ((8, True),) + dense = 248320 * 2 + state = 8 * r._backward_row_state_bytes() + # However the CP split leaves a rank's projected rows, a chunk below the + # fused minimum takes the FP32 fallback: nine buffers of up to 63 rows. + stage = r._checkpoint_head_stage_bytes(3 * 8 * dense, 0, groups, None) + assert stage == 9 * 63 * dense + state + stage = r._checkpoint_head_stage_bytes(3 * 512 * dense, 0, groups, None) + assert stage == 3 * 512 * dense + state + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", "512") + stage = r._checkpoint_head_stage_bytes(3 * 512 * dense, 0, groups, None) + assert stage == 9 * 511 * dense + state From cd7539f5ba00de1b7c3ef7c88e93917388c93bdf Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 14:53:24 +0000 Subject: [PATCH 25/40] Learn memory profiles without the TE workspaces a wave first allocates TE caches its cuBLAS workspaces for the process: the plan whose GEMMs first allocate them peaks higher than any later plan will, yet its reading set the per-token rate that prices every later wave until a warm reading exists. Exclude the workspace growth observed within the plan's peak window from its reading: exact when both workspace sets appear, else one workspace per new set, the least any holds. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 55 +++++++++++++++++++- tests/unit/test_trainer_rank_profile_warm.py | 37 +++++++++++++ 2 files changed, 90 insertions(+), 2 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 672e924b0..1af300f95 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4564,6 +4564,43 @@ def _generic_checkpoint_floor( workspace += self._te_workspace_growth_bytes() return retained, workspace + @staticmethod + def _te_workspace_sets() -> int | None: + """TE's cached cuBLAS workspace sets, or None without TE.""" + try: + from transformer_engine.pytorch.cpp_extensions import gemm + except ImportError: + return None + info = getattr(gemm.get_cublas_workspace, "cache_info", None) + return int(info().currsize) if callable(info) else None + + def _te_workspace_growth_since(self, sets: int | None) -> int: + """TE cuBLAS workspace bytes allocated since ``sets`` sets were cached. + + TE caches one set per kind: the plain GEMMs' one workspace and the + grouped GEMMs' one per cuBLAS stream (userbuffers' only with TP comm + overlap). Growing from none to both is exact; any other growth is + ambiguous and counts one workspace per set, the least any holds. + """ + now = self._te_workspace_sets() + if sets is None or now is None or now <= sets: + return 0 + try: + from transformer_engine.pytorch.cpp_extensions import gemm + + size = int(gemm.get_cublas_workspace_size_bytes()) + streams = int(gemm.tex.get_num_cublas_streams()) + except (AttributeError, ImportError, RuntimeError): + return 0 + overlap = bool( + getattr( + getattr(self.runtime.model[0], "config", None), "tp_comm_overlap", False + ) + ) + if sets == 0 and now == 2 and not overlap: + return (1 + streams) * size + return (now - sets) * size + def _te_workspace_growth_bytes(self) -> int: """Transformer Engine's cuBLAS workspaces, until its GEMMs allocate them. @@ -7047,6 +7084,10 @@ def _run_flat_plan_with_memory_tracking( baseline = int(torch.cuda.memory_allocated(self.device)) torch.cuda.reset_peak_memory_stats(self.device) self._peak_resets = self.__dict__.get("_peak_resets", 0) + 1 + self._te_workspaces_at_reset = ( + self._peak_resets, + self._te_workspace_sets(), + ) else: baseline = None observation = getattr(self, "_planner_observation", None) @@ -7109,11 +7150,21 @@ def _update_peak_memory_profile( observation["peak"] = max(observation["peak"], peak) resets = self.__dict__.get("_peak_resets", 0) self._peak_reading = (resets, peak) + # TE's cuBLAS workspaces this window allocated stay for the process: + # a one-time cost later plans never pay, so it teaches no rate. + at_reset = self.__dict__.get("_te_workspaces_at_reset") + growth = ( + self._te_workspace_growth_since(at_reset[1]) + if at_reset is not None and at_reset[0] == resets + else 0 + ) self._update_memory_profile( plan, - max(0, peak - baseline), + max(0, peak - baseline - growth), retained_bytes=( - None if retained_after is None else max(0, retained_after - baseline) + None + if retained_after is None + else max(0, retained_after - baseline - growth) ), # Only a whole wave: a nested forward resetting the counter during # the yield, or an untracked reset (the counter fell below this diff --git a/tests/unit/test_trainer_rank_profile_warm.py b/tests/unit/test_trainer_rank_profile_warm.py index f1fa1c38e..7248368cc 100644 --- a/tests/unit/test_trainer_rank_profile_warm.py +++ b/tests/unit/test_trainer_rank_profile_warm.py @@ -8,6 +8,7 @@ from dataclasses import replace +import pytest from test_trainer_rank_active_memory import _rank as packed_rank from test_trainer_rank_active_memory import _requests as packed_requests from test_trainer_rank_checkpoint_memory import rank, requests @@ -471,3 +472,39 @@ def test_a_higher_reading_before_the_first_warm_plan_is_kept(): profile = r._memory_profiles[first.signature] assert profile.warm_bytes_per_token == WARM + 300_000 assert _required(r, larger) == int((WARM + 300_000) * larger.packed_tokens * 1.1) + + +@pytest.mark.parametrize( + "before,after,excluded", + [ + (0, 2, 5 * 1000), # Plain and grouped sets: one + four stream workspaces. + (1, 2, 1000), # Either kind: the least a set holds. + (0, 1, 1000), + (2, 2, 0), # Already allocated: later waves never pay it. + ], +) +def test_te_workspaces_a_wave_grows_teach_no_rate(monkeypatch, before, after, excluded): + from transformer_engine.pytorch.cpp_extensions import gemm + + r = rank() + first, _larger = _plans(r) + monkeypatch.setattr(gemm, "get_cublas_workspace_size_bytes", lambda: 1000) + monkeypatch.setattr(gemm.tex, "get_num_cublas_streams", lambda: 4) + sets = {"n": before} + monkeypatch.setattr(r, "_te_workspace_sets", lambda: sets["n"]) + state, _ = _peak_reader(monkeypatch, r) + r._peak_resets = 7 + r._te_workspaces_at_reset = (7, before) + rate = 300 + state["peak"] = first.output_bytes + rate * first.packed_tokens + excluded + sets["n"] = after + r._update_peak_memory_profile(first, 0, 0) + profile = r._memory_profiles[first.signature] + assert profile.bytes_per_token == rate + # A reading outside the window that saw the sets grow keeps them. + del r._memory_profiles[first.signature] + r._te_workspaces_at_reset = (6, before) + r._update_peak_memory_profile(first, 0, 0) + assert r._memory_profiles[first.signature].bytes_per_token == ( + rate + excluded / first.packed_tokens + ) From e88b5f3f1dd8ba0a7e05d9fc80edd6debf66e12c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 15:05:03 +0000 Subject: [PATCH 26/40] Revert TE-free profile readings The TE workspaces a window allocates need not have existed at its peak (an earlier transient can set the maximum before a later GEMM first allocates them), and TE's cache count is process-wide, not per device. Neither is provable from the counters, so the reading keeps them, as before. This reverts commit cd7539f5ba00de1b7c3ef7c88e93917388c93bdf. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 55 +------------------- tests/unit/test_trainer_rank_profile_warm.py | 37 ------------- 2 files changed, 2 insertions(+), 90 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 1af300f95..672e924b0 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4564,43 +4564,6 @@ def _generic_checkpoint_floor( workspace += self._te_workspace_growth_bytes() return retained, workspace - @staticmethod - def _te_workspace_sets() -> int | None: - """TE's cached cuBLAS workspace sets, or None without TE.""" - try: - from transformer_engine.pytorch.cpp_extensions import gemm - except ImportError: - return None - info = getattr(gemm.get_cublas_workspace, "cache_info", None) - return int(info().currsize) if callable(info) else None - - def _te_workspace_growth_since(self, sets: int | None) -> int: - """TE cuBLAS workspace bytes allocated since ``sets`` sets were cached. - - TE caches one set per kind: the plain GEMMs' one workspace and the - grouped GEMMs' one per cuBLAS stream (userbuffers' only with TP comm - overlap). Growing from none to both is exact; any other growth is - ambiguous and counts one workspace per set, the least any holds. - """ - now = self._te_workspace_sets() - if sets is None or now is None or now <= sets: - return 0 - try: - from transformer_engine.pytorch.cpp_extensions import gemm - - size = int(gemm.get_cublas_workspace_size_bytes()) - streams = int(gemm.tex.get_num_cublas_streams()) - except (AttributeError, ImportError, RuntimeError): - return 0 - overlap = bool( - getattr( - getattr(self.runtime.model[0], "config", None), "tp_comm_overlap", False - ) - ) - if sets == 0 and now == 2 and not overlap: - return (1 + streams) * size - return (now - sets) * size - def _te_workspace_growth_bytes(self) -> int: """Transformer Engine's cuBLAS workspaces, until its GEMMs allocate them. @@ -7084,10 +7047,6 @@ def _run_flat_plan_with_memory_tracking( baseline = int(torch.cuda.memory_allocated(self.device)) torch.cuda.reset_peak_memory_stats(self.device) self._peak_resets = self.__dict__.get("_peak_resets", 0) + 1 - self._te_workspaces_at_reset = ( - self._peak_resets, - self._te_workspace_sets(), - ) else: baseline = None observation = getattr(self, "_planner_observation", None) @@ -7150,21 +7109,11 @@ def _update_peak_memory_profile( observation["peak"] = max(observation["peak"], peak) resets = self.__dict__.get("_peak_resets", 0) self._peak_reading = (resets, peak) - # TE's cuBLAS workspaces this window allocated stay for the process: - # a one-time cost later plans never pay, so it teaches no rate. - at_reset = self.__dict__.get("_te_workspaces_at_reset") - growth = ( - self._te_workspace_growth_since(at_reset[1]) - if at_reset is not None and at_reset[0] == resets - else 0 - ) self._update_memory_profile( plan, - max(0, peak - baseline - growth), + max(0, peak - baseline), retained_bytes=( - None - if retained_after is None - else max(0, retained_after - baseline - growth) + None if retained_after is None else max(0, retained_after - baseline) ), # Only a whole wave: a nested forward resetting the counter during # the yield, or an untracked reset (the counter fell below this diff --git a/tests/unit/test_trainer_rank_profile_warm.py b/tests/unit/test_trainer_rank_profile_warm.py index 7248368cc..f1fa1c38e 100644 --- a/tests/unit/test_trainer_rank_profile_warm.py +++ b/tests/unit/test_trainer_rank_profile_warm.py @@ -8,7 +8,6 @@ from dataclasses import replace -import pytest from test_trainer_rank_active_memory import _rank as packed_rank from test_trainer_rank_active_memory import _requests as packed_requests from test_trainer_rank_checkpoint_memory import rank, requests @@ -472,39 +471,3 @@ def test_a_higher_reading_before_the_first_warm_plan_is_kept(): profile = r._memory_profiles[first.signature] assert profile.warm_bytes_per_token == WARM + 300_000 assert _required(r, larger) == int((WARM + 300_000) * larger.packed_tokens * 1.1) - - -@pytest.mark.parametrize( - "before,after,excluded", - [ - (0, 2, 5 * 1000), # Plain and grouped sets: one + four stream workspaces. - (1, 2, 1000), # Either kind: the least a set holds. - (0, 1, 1000), - (2, 2, 0), # Already allocated: later waves never pay it. - ], -) -def test_te_workspaces_a_wave_grows_teach_no_rate(monkeypatch, before, after, excluded): - from transformer_engine.pytorch.cpp_extensions import gemm - - r = rank() - first, _larger = _plans(r) - monkeypatch.setattr(gemm, "get_cublas_workspace_size_bytes", lambda: 1000) - monkeypatch.setattr(gemm.tex, "get_num_cublas_streams", lambda: 4) - sets = {"n": before} - monkeypatch.setattr(r, "_te_workspace_sets", lambda: sets["n"]) - state, _ = _peak_reader(monkeypatch, r) - r._peak_resets = 7 - r._te_workspaces_at_reset = (7, before) - rate = 300 - state["peak"] = first.output_bytes + rate * first.packed_tokens + excluded - sets["n"] = after - r._update_peak_memory_profile(first, 0, 0) - profile = r._memory_profiles[first.signature] - assert profile.bytes_per_token == rate - # A reading outside the window that saw the sets grow keeps them. - del r._memory_profiles[first.signature] - r._te_workspaces_at_reset = (6, before) - r._update_peak_memory_profile(first, 0, 0) - assert r._memory_profiles[first.signature].bytes_per_token == ( - rate + excluded / first.packed_tokens - ) From 8aecea601558e5a2fc685dce78edc3b45a84c769 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 15:09:24 +0000 Subject: [PATCH 27/40] Run a staged plan's fused head statistics strictly A kernel that fails after earlier successes falls back to FP32 statistics silently; marking the failure protects only later plans, not the one already admitted on the fused path's memory. Record each plan's latest staging on it, and have its execution run the fused statistics strictly, through the checkpoint recompute: a failure then raises rather than exceeding the price. Plans priced unstaged keep the silent fallback under the unstaged floor. Several labels per row are not traced either: their gather saves indices and masks per label. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 59 +++++++++++++-- .../test_trainer_rank_head_stage_memory.py | 71 +++++++++++++++++++ 2 files changed, 123 insertions(+), 7 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 672e924b0..9124ee696 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4087,8 +4087,9 @@ def _head_backward_traced( projected chunk. Top-k, logits and hidden-state outputs keep further dense gradients, and the FP32 fallback wider copies; other CP sizes are untraced. The fused statistics must have run in this process and - never fallen back after an error: a failed kernel takes the FP32 path - silently. ``rows`` bounds the first chunk's rows from below, + never fallen back after an error; a plan priced on them then runs them + strictly (``_plan_head_backward_traced``), so a later failure raises + rather than taking the wider FP32 path. ``rows`` bounds the first chunk's rows from below, ``upper_rows`` from above: None when they straddle the threshold. A CP rank projecting fewer rows falls back on its own; the head stage prices that (``_HEAD_FALLBACK_BUFFERS``). @@ -4097,6 +4098,8 @@ def _head_backward_traced( not requests or any( request.target_tokens is None + # Several labels per row save gather indices and masks per label. + or request.target_tokens.shape != request.input_tokens.shape or request.top_k is not None or request.logits or request.hidden_states @@ -4119,7 +4122,11 @@ def _head_backward_traced( return None def _plan_head_backward_traced(self, plan: _FlatForwardPlan) -> bool: - """Every gradient group's head backward is traced (``_head_backward_traced``).""" + """Every gradient group's head backward is traced (``_head_backward_traced``). + + Records the answer on the plan: its execution runs the fused statistics + strictly exactly when its latest price staged the head. + """ traced = [ self._head_backward_traced( requests, @@ -4133,7 +4140,9 @@ def _plan_head_backward_traced(self, plan: _FlatForwardPlan) -> bool: if group.grad_enabled for requests in (tuple(item.request for item in group.items),) ] - return bool(traced) and all(state is True for state in traced) + staged = bool(traced) and all(state is True for state in traced) + object.__setattr__(plan, "_head_staged", staged) + return staged def _plan_head_workspace_bytes(self, plan: _FlatForwardPlan) -> int: peak = 0 @@ -7587,6 +7596,17 @@ def _telemetry_signature(cls, plan: _AnyForwardPlan) -> dict[str, object]: } def _execute_flat_plan(self, plan: _FlatForwardPlan) -> list[AnyForwardOutput]: + # A head priced as staged must not silently widen (_head_backward_traced). + previous_strict = self.__dict__.get("_head_statistics_strict", False) + self._head_statistics_strict = bool(getattr(plan, "_head_staged", False)) + try: + return self._execute_flat_plan_groups(plan) + finally: + self._head_statistics_strict = previous_strict + + def _execute_flat_plan_groups( + self, plan: _FlatForwardPlan + ) -> list[AnyForwardOutput]: outputs = [ ForwardOutput(None, None, None, None, checkpoint, no_grad) for checkpoint, no_grad in plan.output_metadata @@ -9537,6 +9557,7 @@ def _project_vocab_parallel( from torch.utils.checkpoint import checkpoint model = _language_model(self.runtime.model[0]) + strict_statistics = bool(self.__dict__.get("_head_statistics_strict", False)) max_top_k = max((int(item.request.top_k or 0) for item in items), default=0) need_log_z = any( item.labels is not None or item.request.top_k is not None for item in items @@ -9554,6 +9575,8 @@ def _project_vocab_parallel( output_weight=output_weight, need_log_z=need_log_z, max_top_k=max_top_k, + # Captured now: the backward recompute keeps this plan's mode. + strict_statistics=strict_statistics, use_reentrant=False, ) logit_start, logit_end = logit_bounds[chunk_index : chunk_index + 2] @@ -9634,6 +9657,7 @@ def _local_head_stats( output_weight: torch.Tensor | None, need_log_z: bool, max_top_k: int, + strict_statistics: bool = False, ) -> tuple[ torch.Tensor, torch.Tensor | None, @@ -9647,11 +9671,17 @@ def _local_head_stats( log_z: torch.Tensor | None = None local_topk: tuple[torch.Tensor, torch.Tensor] | None = None if need_log_z: - topk_stats = _try_triton_local_topk_stats(local_logits, k=max_top_k) + topk_stats = _try_triton_local_topk_stats( + local_logits, k=max_top_k, strict=strict_statistics + ) logsumexp_stats = ( cast( tuple[torch.Tensor, torch.Tensor] | None, - _try_triton_stats("local_logsumexp_stats", local_logits), + _try_triton_stats( + "local_logsumexp_stats", + local_logits, + strict=strict_statistics, + ), ) if topk_stats is None else None @@ -10596,6 +10626,7 @@ def _try_triton_local_topk_stats( local_logits: torch.Tensor, *, k: int, + strict: bool = False, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] | None: if k <= 0 or k > int( os.environ.get("ART_TRAINER_RANK_TRITON_FUSED_TOPK_MAX", "10") @@ -10606,6 +10637,7 @@ def _try_triton_local_topk_stats( _try_triton_stats( "local_topk_stats", local_logits, + strict=strict, k=min(k, int(local_logits.shape[1])), ), ) @@ -10614,8 +10646,16 @@ def _try_triton_local_topk_stats( def _try_triton_stats( name: str, local_logits: torch.Tensor, + *, + strict: bool = False, **kwargs: object, ) -> object | None: + """The fused statistics, or None for the FP32 fallback. + + ``strict``: a plan whose price relied on them raises instead of falling + back after an error (``_head_backward_traced``). Too few rows still fall + back: the head stage prices that (``_HEAD_FALLBACK_BUFFERS``). + """ if not local_logits.is_cuda: return None if os.environ.get("ART_TRAINER_RANK_TRITON_TOPK", "1").lower() in { @@ -10629,8 +10669,13 @@ def _try_triton_stats( from art.trainer_rank import topk result = getattr(topk, name)(local_logits, **kwargs) - except Exception: + except Exception as error: _TRITON_STATS_STATE["failed"] = True + if strict: + raise RuntimeError( + "Fused head statistics failed in a plan admitted on their " + "memory; the FP32 fallback would exceed its price" + ) from error if os.environ.get("ART_TRAINER_RANK_TRITON_TOPK", "1").lower() == "strict": raise return None diff --git a/tests/unit/test_trainer_rank_head_stage_memory.py b/tests/unit/test_trainer_rank_head_stage_memory.py index 4c76f6e86..f14118842 100644 --- a/tests/unit/test_trainer_rank_head_stage_memory.py +++ b/tests/unit/test_trainer_rank_head_stage_memory.py @@ -337,3 +337,74 @@ def test_head_stage_covers_a_small_chunks_fp32_fallback(monkeypatch): monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", "512") stage = r._checkpoint_head_stage_bytes(3 * 512 * dense, 0, groups, None) assert stage == 9 * 511 * dense + state + + +def test_a_staged_plan_runs_its_fused_statistics_strictly(monkeypatch): + from types import SimpleNamespace + + from art.trainer_rank import _impl, topk + + state = {"succeeded": True, "failed": False} + monkeypatch.setattr(_impl, "_TRITON_STATS_STATE", state) + + def fail(*args, **kwargs): + raise RuntimeError("kernel launch failed") + + monkeypatch.setattr(topk, "local_logsumexp_stats", fail) + chunk = SimpleNamespace(is_cuda=True, shape=(512, 248320)) + # Unstaged: the FP32 fallback, and no later plan stages again. + assert _impl._try_triton_stats("local_logsumexp_stats", chunk) is None + assert state["failed"] is True + # Staged: the plan's price assumed the fused path, so it raises instead. + with pytest.raises(RuntimeError, match="admitted on their memory"): + _impl._try_triton_stats("local_logsumexp_stats", chunk, strict=True) + # Too few rows is a predictable fallback the head stage prices. + small = SimpleNamespace(is_cuda=True, shape=(63, 248320)) + assert _impl._try_triton_stats("local_logsumexp_stats", small, strict=True) is None + + +def test_execution_binds_the_staging_its_latest_price_used(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + from art.trainer_rank import _impl + + monkeypatch.setattr( + _impl, "_TRITON_STATS_STATE", {"succeeded": True, "failed": False} + ) + r = head_rank() + plan = r._plan_flat_forward([request(512, grad=True)]) + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + seen = [] + monkeypatch.setattr( + r, + "_execute_flat_plan_groups", + lambda plan: seen.append(r._head_statistics_strict) or [], + ) + assert r._plan_head_backward_traced(plan) is True + r._execute_flat_plan(plan) + # A later price that no longer stages (a kernel failed since) unbinds it. + _impl._TRITON_STATS_STATE["failed"] = True + assert r._plan_head_backward_traced(plan) is False + r._execute_flat_plan(plan) + assert seen == [True, False] + assert r._head_statistics_strict is False + + +def test_several_labels_per_row_are_not_traced(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + from art.trainer_rank import _impl + + monkeypatch.setattr( + _impl, "_TRITON_STATS_STATE", {"succeeded": True, "failed": False} + ) + r = head_rank() + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + single = request(512, grad=True) + several = replace( + single, target_tokens=single.target_tokens.unsqueeze(1).repeat(1, 4) + ) + assert r._head_backward_traced([single], 512) is True + assert r._head_backward_traced([several], 512) is False From b0c9cfac37d19e27ac87c0a7fecd4f09ec9d8bd6 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 15:16:11 +0000 Subject: [PATCH 28/40] Prove the target-only fused kernel itself before staging A top-k or no-grad success compiles other specializations: record which fused statistics kernels have run, and stage target-only heads only once the logsumexp kernel has. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 10 +++++---- .../test_trainer_rank_head_stage_memory.py | 21 +++++++++++++------ 2 files changed, 21 insertions(+), 10 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 9124ee696..daad151ac 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -142,9 +142,9 @@ class TopK: # BF16 logits-sized buffers of that chunk (about 18 bytes per logit); derived # from the code, not traced. _HEAD_FALLBACK_BUFFERS = 9 -# Whether the fused head statistics have run in this process, and whether any +# Which fused head statistics kernels have run in this process, and whether any # call fell back to FP32 after an error: staging trusts only a proven path. -_TRITON_STATS_STATE = {"succeeded": False, "failed": False} +_TRITON_STATS_STATE: dict[str, Any] = {"succeeded": set(), "failed": False} _PLANNER_REFINEMENT_BUDGET = 2_000 _LAYOUT_SELECTION_CACHE_LIMIT = 64 @@ -4107,7 +4107,9 @@ def _head_backward_traced( ) or os.environ.get("ART_TRAINER_RANK_TRITON_TOPK", "1").lower() in {"0", "false"} - or not _TRITON_STATS_STATE["succeeded"] + # Target-only heads run the logsumexp kernel; another kernel's + # success does not prove it. + or "local_logsumexp_stats" not in _TRITON_STATS_STATE["succeeded"] or _TRITON_STATS_STATE["failed"] or self._topology_key()[2] != 2 or not self._head_workspace_bytes(1) @@ -10679,7 +10681,7 @@ def _try_triton_stats( if os.environ.get("ART_TRAINER_RANK_TRITON_TOPK", "1").lower() == "strict": raise return None - _TRITON_STATS_STATE["succeeded"] = True + _TRITON_STATS_STATE["succeeded"].add(name) return result diff --git a/tests/unit/test_trainer_rank_head_stage_memory.py b/tests/unit/test_trainer_rank_head_stage_memory.py index f14118842..2cabd09a8 100644 --- a/tests/unit/test_trainer_rank_head_stage_memory.py +++ b/tests/unit/test_trainer_rank_head_stage_memory.py @@ -245,10 +245,13 @@ def test_traced_head_backward_is_target_only_on_the_traced_path(monkeypatch): monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) target = request(512, grad=True) # Until the fused statistics have run in this process, nothing is traced. - state = {"succeeded": False, "failed": False} + state = {"succeeded": set(), "failed": False} monkeypatch.setattr(_impl, "_TRITON_STATS_STATE", state) assert r._head_backward_traced([target], 512) is False - state["succeeded"] = True + # The top-k kernel's success does not prove the target-only one. + state["succeeded"].add("local_topk_stats") + assert r._head_backward_traced([target], 512) is False + state["succeeded"].add("local_logsumexp_stats") assert r._head_backward_traced([target], 512) is True # Top-k, logits and hidden-state outputs keep further dense gradients. for extra in ({"top_k": 2}, {"logits": True}, {"hidden_states": True}): @@ -304,7 +307,9 @@ def test_plan_stages_only_traced_gradient_heads(monkeypatch): from art.trainer_rank import _impl monkeypatch.setattr( - _impl, "_TRITON_STATS_STATE", {"succeeded": True, "failed": False} + _impl, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, ) r = head_rank() target = request(512, grad=True) @@ -344,7 +349,7 @@ def test_a_staged_plan_runs_its_fused_statistics_strictly(monkeypatch): from art.trainer_rank import _impl, topk - state = {"succeeded": True, "failed": False} + state = {"succeeded": {"local_logsumexp_stats"}, "failed": False} monkeypatch.setattr(_impl, "_TRITON_STATS_STATE", state) def fail(*args, **kwargs): @@ -370,7 +375,9 @@ def test_execution_binds_the_staging_its_latest_price_used(monkeypatch): from art.trainer_rank import _impl monkeypatch.setattr( - _impl, "_TRITON_STATS_STATE", {"succeeded": True, "failed": False} + _impl, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, ) r = head_rank() plan = r._plan_flat_forward([request(512, grad=True)]) @@ -398,7 +405,9 @@ def test_several_labels_per_row_are_not_traced(monkeypatch): from art.trainer_rank import _impl monkeypatch.setattr( - _impl, "_TRITON_STATS_STATE", {"succeeded": True, "failed": False} + _impl, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, ) r = head_rank() monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) From 35850f7449075e8d55d2b2e9ca044a10911c144e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 15:33:48 +0000 Subject: [PATCH 29/40] Bind strict head statistics per thread, gradient group and staged admission A shared instance flag let a concurrent execution flip another's mode, a diagnostic re-price after admission could clear it, eligibility marked plans whose price never staged, and no-grad groups ran strict too. Mark a plan only when its price stages the head, never unmark it, and scope the mode to each gradient group's projection in a context variable. Count one label per input token in either accepted layout as a single label. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 60 ++++++++++++------- .../test_trainer_rank_head_stage_memory.py | 60 +++++++++++++++---- 2 files changed, 88 insertions(+), 32 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index daad151ac..467a3289e 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -145,6 +145,11 @@ class TopK: # Which fused head statistics kernels have run in this process, and whether any # call fell back to FP32 after an error: staging trusts only a proven path. _TRITON_STATS_STATE: dict[str, Any] = {"succeeded": set(), "failed": False} +# Set while a gradient group of a plan priced with a staged head projects it; +# per thread, and captured into each chunk's checkpoint for its recompute. +_HEAD_STATISTICS_STRICT: ContextVar[bool] = ContextVar( + "trainer_rank_head_statistics_strict", default=False +) _PLANNER_REFINEMENT_BUDGET = 2_000 _LAYOUT_SELECTION_CACHE_LIMIT = 64 @@ -4092,14 +4097,17 @@ def _head_backward_traced( rather than taking the wider FP32 path. ``rows`` bounds the first chunk's rows from below, ``upper_rows`` from above: None when they straddle the threshold. A CP rank projecting fewer rows falls back on its own; the head stage - prices that (``_HEAD_FALLBACK_BUFFERS``). + prices that (``_HEAD_FALLBACK_BUFFERS``). ``ART_TRAINER_RANK_TRITON_*`` + settings are process-wide: changing them while a plan is in flight is + unsupported (a staged chunk could then fall back unpriced). """ if ( not requests or any( request.target_tokens is None - # Several labels per row save gather indices and masks per label. - or request.target_tokens.shape != request.input_tokens.shape + # Several labels per row save gather indices and masks per label; + # one label per input token, in either accepted layout, does not. + or request.target_tokens.numel() != request.input_tokens.numel() or request.top_k is not None or request.logits or request.hidden_states @@ -4126,8 +4134,16 @@ def _head_backward_traced( def _plan_head_backward_traced(self, plan: _FlatForwardPlan) -> bool: """Every gradient group's head backward is traced (``_head_backward_traced``). - Records the answer on the plan: its execution runs the fused statistics - strictly exactly when its latest price staged the head. + When that stages the plan's head price (``_checkpoint_head_stage_bytes`` + applies), marks the plan: its gradient groups then run the fused + statistics strictly (``_execute_flat_plan``). The mark stays once set: + a later re-pricing cannot weaken the admission that relied on it. + Every plan admitted on a staged price has been priced here: staging + needs CP2, where the cheap width estimate declines and admission + prices the materialized plan (``_estimate_flat_forward``). Strictness + trades the FP32 fallback's availability for memory safety: a fused + kernel error then fails the wave, and a CP peer waits in its next + collective like after a rank-local OOM. """ traced = [ self._head_backward_traced( @@ -4142,9 +4158,16 @@ def _plan_head_backward_traced(self, plan: _FlatForwardPlan) -> bool: if group.grad_enabled for requests in (tuple(item.request for item in group.items),) ] - staged = bool(traced) and all(state is True for state in traced) - object.__setattr__(plan, "_head_staged", staged) - return staged + eligible = bool(traced) and all(state is True for state in traced) + if ( + eligible + and self._plan_head_workspace_bytes(plan) + and self._checkpoint_gradient_covered( + self._plan_group_rows(plan), tuple(g.slot_ref for g in plan.groups) + ) + ): + object.__setattr__(plan, "_head_staged", True) + return eligible def _plan_head_workspace_bytes(self, plan: _FlatForwardPlan) -> int: peak = 0 @@ -7599,16 +7622,7 @@ def _telemetry_signature(cls, plan: _AnyForwardPlan) -> dict[str, object]: def _execute_flat_plan(self, plan: _FlatForwardPlan) -> list[AnyForwardOutput]: # A head priced as staged must not silently widen (_head_backward_traced). - previous_strict = self.__dict__.get("_head_statistics_strict", False) - self._head_statistics_strict = bool(getattr(plan, "_head_staged", False)) - try: - return self._execute_flat_plan_groups(plan) - finally: - self._head_statistics_strict = previous_strict - - def _execute_flat_plan_groups( - self, plan: _FlatForwardPlan - ) -> list[AnyForwardOutput]: + staged = bool(getattr(plan, "_head_staged", False)) outputs = [ ForwardOutput(None, None, None, None, checkpoint, no_grad) for checkpoint, no_grad in plan.output_metadata @@ -7630,7 +7644,13 @@ def _execute_flat_plan_groups( with torch.set_grad_enabled(group.grad_enabled): with use_lora_slot(group.slot_ref): prepared = self._prepare_packed_forward(group.packed) - item_outputs = self._forward_packed(group.items, prepared) + strict = _HEAD_STATISTICS_STRICT.set( + staged and group.grad_enabled + ) + try: + item_outputs = self._forward_packed(group.items, prepared) + finally: + _HEAD_STATISTICS_STRICT.reset(strict) item_outputs = [ replace( output, @@ -9559,7 +9579,7 @@ def _project_vocab_parallel( from torch.utils.checkpoint import checkpoint model = _language_model(self.runtime.model[0]) - strict_statistics = bool(self.__dict__.get("_head_statistics_strict", False)) + strict_statistics = _HEAD_STATISTICS_STRICT.get() max_top_k = max((int(item.request.top_k or 0) for item in items), default=0) need_log_z = any( item.labels is not None or item.request.top_k is not None for item in items diff --git a/tests/unit/test_trainer_rank_head_stage_memory.py b/tests/unit/test_trainer_rank_head_stage_memory.py index 2cabd09a8..243325310 100644 --- a/tests/unit/test_trainer_rank_head_stage_memory.py +++ b/tests/unit/test_trainer_rank_head_stage_memory.py @@ -368,11 +368,11 @@ def fail(*args, **kwargs): assert _impl._try_triton_stats("local_logsumexp_stats", small, strict=True) is None -def test_execution_binds_the_staging_its_latest_price_used(monkeypatch): +def test_execution_binds_strictness_to_the_staged_admission(monkeypatch): from test_trainer_rank_head_memory import rank as head_rank from test_trainer_rank_head_memory import request - from art.trainer_rank import _impl + from art.trainer_rank import ForwardOutput, _impl monkeypatch.setattr( _impl, @@ -380,22 +380,55 @@ def test_execution_binds_the_staging_its_latest_price_used(monkeypatch): {"succeeded": {"local_logsumexp_stats"}, "failed": False}, ) r = head_rank() - plan = r._plan_flat_forward([request(512, grad=True)]) + monkeypatch.setattr(r, "_moe_recompute_covered_for", lambda ref: True) + plan = r._plan_flat_forward([request(512, grad=True), request(16, hidden=True)]) monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + monkeypatch.setattr(r, "_topology", lambda: None) + monkeypatch.setattr(r, "_validate_hybridep_topology", lambda: None) + monkeypatch.setattr(r, "_configure_hybridep", lambda *a, **k: None) + monkeypatch.setattr(r, "_prepare_packed_forward", lambda packed: None) seen = [] - monkeypatch.setattr( - r, - "_execute_flat_plan_groups", - lambda plan: seen.append(r._head_statistics_strict) or [], - ) + + def forward_packed(items, prepared): + seen.append(_impl._HEAD_STATISTICS_STRICT.get()) + return [ForwardOutput(None, None, None, None)] * len(items) + + monkeypatch.setattr(r, "_forward_packed", forward_packed) + r._execute_flat_plan(plan) + assert seen == [False, False] # Never priced staged. assert r._plan_head_backward_traced(plan) is True + assert plan._head_staged is True + seen.clear() r._execute_flat_plan(plan) - # A later price that no longer stages (a kernel failed since) unbinds it. + # Only the gradient group runs strictly; nothing leaks past execution. + grad_first = [group.grad_enabled for group in plan.groups] + assert seen == grad_first + assert _impl._HEAD_STATISTICS_STRICT.get() is False + # A later price that no longer stages (a kernel failed since) cannot + # weaken the admission that relied on it. _impl._TRITON_STATS_STATE["failed"] = True assert r._plan_head_backward_traced(plan) is False - r._execute_flat_plan(plan) - assert seen == [True, False] - assert r._head_statistics_strict is False + assert plan._head_staged is True + + +def test_an_eligible_head_the_price_does_not_stage_is_not_strict(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + from art.trainer_rank import _impl + + monkeypatch.setattr( + _impl, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, + ) + r = head_rank() + plan = r._plan_flat_forward([request(512, grad=True)]) + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + # The recompute is not covered: the head keeps the unstaged price. + monkeypatch.setattr(r, "_checkpoint_gradient_covered", lambda *a, **k: False) + assert r._plan_head_backward_traced(plan) is True + assert getattr(plan, "_head_staged", False) is False def test_several_labels_per_row_are_not_traced(monkeypatch): @@ -417,3 +450,6 @@ def test_several_labels_per_row_are_not_traced(monkeypatch): ) assert r._head_backward_traced([single], 512) is True assert r._head_backward_traced([several], 512) is False + # One label per token over a leading batch axis is still one per row. + batched = replace(single, input_tokens=single.input_tokens.unsqueeze(0)) + assert r._head_backward_traced([batched], 512) is True From be1f15df12409d156606457cb09ad39a5f13571c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 03:51:09 +0000 Subject: [PATCH 30/40] Replay the adapter-gradient floor from frozen slot facts Grouped planner replay (#1028) recomputes each subforward's cost from primitive runtime facts, with group indices standing in for slot refs. The adapter-gradient floor reads the gradient slot's unallocated LoRA gradients from the live model, so replay could neither resolve the slot nor reproduce the term. Capture each gradient group's slot kind and name and its pending gradient bytes per decoder layer with the selection (runtime facts version 2), and have ReplayRank answer the floor's slot and pending-gradient readers from those facts. The capture's stock-estimator check now covers the floor's readers too. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 92 ++++++++++++++++++- tests/unit/test_grouped_planner_replay.py | 55 ++++++++++- .../unit/test_planner_replay_owner_budget.py | 2 +- 3 files changed, 143 insertions(+), 6 deletions(-) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index f0da8bbe1..15786f0b0 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -11,7 +11,7 @@ from dataclasses import asdict import json from types import MethodType, SimpleNamespace -from typing import Any +from typing import Any, NamedTuple from . import _gdn_memory, _impl, _memory @@ -21,6 +21,12 @@ _MAX_INPUT_VALUES = 1_000_000 +# Estimators bound on TrainerRank as plain functions, not methods. +_STATIC_ESTIMATORS = frozenset( + {"_split_required_memory", "_gradient_slots", "_adapter_gradient_walk"} +) + + _REFUSALS = frozenset( { "runtime_group_inventory_over_limit", @@ -96,12 +102,17 @@ def capture(rank: Any, plan: Any) -> dict[str, Any]: "_physical_tokens", "_plan_group_rows", "_plan_retained_tokens", + "_gradient_slots", + "_pending_adapter_gradient_bytes", + "_checkpoint_gradient_groups", + "_checkpoint_adapter_gradient_bytes", + "_adapter_gradient_walk", ): method = getattr(rank, name) expected = getattr(_impl.TrainerRank, name) supported = ( method is expected - if name == "_split_required_memory" # The sole static estimator. + if name in _STATIC_ESTIMATORS else type(method) is MethodType and method.__self__ is rank and method.__func__ is expected @@ -214,6 +225,18 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: if backwards else 0 ) + adapter = None + if group.grad_enabled and getattr(group.slot_ref, "name", None) is not None: + # The live estimator reads this slot's unallocated gradient bytes + # per decoder layer; freeze them with the selection. + kind = getattr(group.slot_ref, "kind", None) + if kind is not None and (type(kind) is not str or len(kind) > 64): + raise ValueError("runtime_slot_identity_unsupported") + pending = rank._pending_adapter_gradient_bytes((group.slot_ref,)) + if len(pending) > 1025: + raise ValueError("runtime_shape_inventory_over_limit") + reserve(128 + 12 * len(name) + 24 * len(pending)) + adapter = {"kind": kind, "name": name, "pending": [int(v) for v in pending]} model = _gdn_memory.model_shapes(rank, group.slot_ref) if has_grad else None if model is not None: if len(model[1]) > 1024: @@ -231,6 +254,7 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: "gradient": terms(True, group.slot_ref), "head_rows": projected, "head_target_rows": target_rows, + "adapter": adapter, "gdn": None if model is None else { @@ -253,7 +277,7 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: } ) facts = { - "version": 1, + "version": 2, "checkpoint_layers": _memory._checkpoint_layers( rank, rank._plan_group_rows(plan) ), @@ -290,7 +314,7 @@ def integer(value: Any, *, minimum: int = 0) -> None: "groups", }, ) - if type(facts["version"]) is not int or facts["version"] != 1: + if type(facts["version"]) is not int or facts["version"] != 2: raise ValueError("unsupported runtime facts version") for key in ( "checkpoint_layers", @@ -319,6 +343,7 @@ def integer(value: Any, *, minimum: int = 0) -> None: "gradient", "head_rows", "head_target_rows", + "adapter", "gdn", }, ) @@ -363,6 +388,22 @@ def integer(value: Any, *, minimum: int = 0) -> None: raise ValueError("invalid MoE stage") for value in stage: integer(value) + adapter = group["adapter"] + if adapter is not None: + fields(adapter, {"kind", "name", "pending"}) + if ( + not group["grad"] + or (adapter["kind"] is not None and type(adapter["kind"]) is not str) + or len(adapter["kind"] or "") > 64 + or type(adapter["name"]) is not str + or len(adapter["name"]) > 4096 + or type(adapter["pending"]) is not list + or len(adapter["pending"]) > 1025 + ): + raise ValueError("invalid adapter gradient facts") + reserve(128 + 12 * len(adapter["name"]) + 24 * len(adapter["pending"])) + for value in adapter["pending"]: + integer(value) gdn = group["gdn"] if gdn is not None: fields(gdn, {"layers", "shapes", "segments"}) @@ -420,11 +461,54 @@ def count(value: Any, depth: int = 0) -> None: count(value) +class _ReplaySlot(NamedTuple): + """A frozen adapter slot identity (the live LoRASlotRef's kind and name).""" + + kind: str | None + name: str + + class ReplayRank(_impl.TrainerRank): """The real estimator with runtime metadata readers replaced by frozen facts.""" _facts: dict[str, Any] | None = None + def _replay_slots(self, slot_refs: Any) -> Any: + # Replay passes each group's index; map it to that group's frozen slot. + if self._facts is None or slot_refs is None: + return slot_refs + groups = self._facts["groups"] + return tuple( + None + if (adapter := groups[index]["adapter"]) is None + else _ReplaySlot(adapter["kind"], adapter["name"]) + for index in slot_refs + ) + + def _gradient_slots(self, group_rows: Any, slot_refs: Any) -> Any: + return _memory._gradient_slots(group_rows, self._replay_slots(slot_refs)) + + def _checkpoint_gradient_groups(self, group_rows: Any, slot_refs: Any) -> Any: + return _memory._checkpoint_gradient_groups( + self, group_rows, self._replay_slots(slot_refs) + ) + + def _pending_adapter_gradient_bytes(self, refs: Any) -> tuple[int, ...]: + if self._facts is None: + return _memory._pending_adapter_gradient_bytes(self, refs) + refs = tuple(dict.fromkeys(refs)) + if not refs: + return () + if len(refs) != 1: + raise ValueError("replayed adapter gradients are frozen per slot") + for group in self._facts["groups"]: + adapter = group["adapter"] + if adapter is not None and (adapter["kind"], adapter["name"]) == tuple( + refs[0] + ): + return tuple(adapter["pending"]) + return () + def _head_workspace_bytes(self, rows: int) -> int: assert self._facts is not None return _memory._dense_head_bytes(self._facts["head_vocabulary"], rows) diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 575fafa54..5f63d0084 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -171,6 +171,59 @@ def test_selected_slot_terms_are_replayed_and_frozen(layer, tmp_path): ) +def test_selected_adapter_gradients_are_replayed_and_frozen(layer, tmp_path): + from test_trainer_rank_adapter_gradient_memory import lora, parameter + from test_trainer_rank_converted_memory import weights + from test_trainer_rank_pending_memory import rank_with_moe + from test_trainer_rank_slot_memory import load_slot + from test_trainer_rank_slot_memory import request as slot_request + + rank, _ = rank_with_moe(weights(layer, 8)) + load_slot(rank, "small", 1) + load_slot(rank, "large", 64) + # Unallocated slot gradients the recompute backward will allocate: small, + # distinct per-layer sizes (the fixture has 40 layers; keep this bounded). + layers = tr._language_model(rank.runtime.model[0]).decoder.layers + assert len(layers) <= 64 + sizes = [16 * (index + 1) for index in range(len(layers))] + assert sum(sizes) * 2 <= 66_560 # BF16 bytes, checked before allocating. + params = [] + for size, block in zip(sizes, layers, strict=True): + params.append(parameter(size)) + block.add_module("adapter", lora(large=[params[-1]])) + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + plan = rank._plan_flat_forward( + [slot_request("small", rows=2), slot_request("large", rows=65, grad=True)], + ensure_slots=False, + ) + original, costs = emitted(rank, plan, tmp_path) + assert costs[0].checkpoint_adapter_gradient > 0 + groups = original["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "groups" + ] + assert groups[0]["adapter"] is None + assert groups[1]["adapter"]["name"] == "large" and any( + groups[1]["adapter"]["pending"] + ) + actual = reports.replay(original) + assert actual["aggregate"]["matches"] + assert all(item["matches"] for item in actual["estimates"]) + # Gradients allocated after selection cannot change the replayed answer. + for param in params: + param.grad = torch.zeros_like(param) + assert reports.replay(original) == actual + changed = deepcopy(original) + changed["replay"]["memory_replay"]["estimates"][0]["runtime_facts"]["groups"][1][ + "adapter" + ]["pending"][0] += 10**12 + result = reports.replay(changed) + assert not result["estimates"][0]["matches"] + assert ( + result["estimates"][0]["required_bytes"] + > actual["estimates"][0]["required_bytes"] + ) + + @pytest.mark.parametrize( "change", ["version", "group", "layout", "gdn_segment", "budget"] ) @@ -184,7 +237,7 @@ def test_fact_validation_rejects_inconsistent_or_unbounded_input( ) facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] if change == "version": - facts["version"] = 2 + facts["version"] += 1 elif change == "group": facts["groups"][0]["rows"] += 1 elif change == "layout": diff --git a/tests/unit/test_planner_replay_owner_budget.py b/tests/unit/test_planner_replay_owner_budget.py index 5292cf743..248683dd2 100644 --- a/tests/unit/test_planner_replay_owner_budget.py +++ b/tests/unit/test_planner_replay_owner_budget.py @@ -47,7 +47,7 @@ def test_shared_inventory_preflight_precedes_layout_construction( elif case == "invalid_facts": report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ "version" - ] = 2 + ] += 1 reason = "unsupported runtime facts version" else: monkeypatch.setattr(runtime, "_MAX_INPUT_VALUES", 10) From b8a2fd183edabe1dae85a4e93e0e548238d2339f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 03:55:54 +0000 Subject: [PATCH 31/40] Narrow the replay capture's slot name before sizing it Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index 15786f0b0..92a746b66 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -226,7 +226,7 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: else 0 ) adapter = None - if group.grad_enabled and getattr(group.slot_ref, "name", None) is not None: + if group.grad_enabled and name is not None: # The live estimator reads this slot's unallocated gradient bytes # per decoder layer; freeze them with the selection. kind = getattr(group.slot_ref, "kind", None) From 6b1b918aabaa3e3a0a2646057bcd0afe2b2aa083 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:12:59 +0000 Subject: [PATCH 32/40] Bound replay's recorded layer count and tighten adapter facts Replay sizes per-layer boundary tuples from the report's recorded num_layers, which only capture bounded; refuse counts outside capture's 1024-layer limit before any estimator runs. A slot without a kind (a megatron-less reference) has no pending gradients, so reject kindless facts that claim some. Test forged adapter facts, the layer bound, and an instance-overridden pending-gradient reader. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_misses.py | 3 ++ src/art/trainer_rank/_planner_replay.py | 5 ++- tests/unit/test_grouped_planner_replay.py | 55 ++++++++++++++++++++++- 3 files changed, 61 insertions(+), 2 deletions(-) diff --git a/src/art/trainer_rank/_planner_misses.py b/src/art/trainer_rank/_planner_misses.py index bb7e8297c..b1df930f1 100644 --- a/src/art/trainer_rank/_planner_misses.py +++ b/src/art/trainer_rank/_planner_misses.py @@ -600,6 +600,9 @@ def replay( raise ValueError( "incomplete replay: immutable rank fields differ (including MoE stages)" ) + layers = values["num_layers"] + if type(layers) is not int or not 0 < layers <= _planner_replay.MAX_LAYERS: + raise ValueError("incomplete replay: recorded layer count out of bounds") rank = _planner_replay.ReplayRank.__new__(_planner_replay.ReplayRank) for name in _RANK_FIELDS - {"one_layer_recompute"}: setattr(rank, "_" + name, values[name]) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index 92a746b66..45c7acc22 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -17,6 +17,8 @@ _MAX_BYTES = 262144 _MAX_GROUPS = 1024 +# Replay sizes per-layer tuples from this recorded rank field; bound it as capture does. +MAX_LAYERS = 1024 _MAX_SEGMENTS = 4096 _MAX_INPUT_VALUES = 1_000_000 @@ -74,7 +76,7 @@ def reserve(size: int) -> None: def capture(rank: Any, plan: Any) -> dict[str, Any]: - if rank._num_layers > 1024 or not 0 < len(plan.groups) <= _MAX_GROUPS: + if rank._num_layers > MAX_LAYERS or not 0 < len(plan.groups) <= _MAX_GROUPS: raise ValueError("runtime_group_inventory_over_limit") if ( getattr(rank.runtime.provider, "expert_model_parallel_size", 1) > 1 @@ -393,6 +395,7 @@ def integer(value: Any, *, minimum: int = 0) -> None: fields(adapter, {"kind", "name", "pending"}) if ( not group["grad"] + or (adapter["kind"] is None and any(adapter["pending"] or ())) or (adapter["kind"] is not None and type(adapter["kind"]) is not str) or len(adapter["kind"] or "") > 64 or type(adapter["name"]) is not str diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 5f63d0084..29d384df6 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -171,7 +171,8 @@ def test_selected_slot_terms_are_replayed_and_frozen(layer, tmp_path): ) -def test_selected_adapter_gradients_are_replayed_and_frozen(layer, tmp_path): +def adapter_report(layer, tmp_path): + """A grouped report whose gradient slot has pending adapter gradients.""" from test_trainer_rank_adapter_gradient_memory import lora, parameter from test_trainer_rank_converted_memory import weights from test_trainer_rank_pending_memory import rank_with_moe @@ -197,6 +198,11 @@ def test_selected_adapter_gradients_are_replayed_and_frozen(layer, tmp_path): ensure_slots=False, ) original, costs = emitted(rank, plan, tmp_path) + return original, costs, params + + +def test_selected_adapter_gradients_are_replayed_and_frozen(layer, tmp_path): + original, costs, params = adapter_report(layer, tmp_path) assert costs[0].checkpoint_adapter_gradient > 0 groups = original["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ "groups" @@ -224,6 +230,53 @@ def test_selected_adapter_gradients_are_replayed_and_frozen(layer, tmp_path): ) +@pytest.mark.parametrize( + "change", ["gradient", "length", "value", "kind_length", "kindless_pending"] +) +def test_adapter_fact_validation_rejects_forged_input(change, layer, tmp_path): + report, _, _ = adapter_report(layer, tmp_path) + group = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "groups" + ][1] + adapter = group["adapter"] + if change == "gradient": + group["grad"] = False + elif change == "length": + adapter["pending"] = [0] * 1026 + elif change == "value": + adapter["pending"][0] = 1.5 + elif change == "kind_length": + adapter["kind"] = "k" * 65 + else: + adapter["kind"] = None + with pytest.raises(ValueError): + reports.replay(report) + + +@pytest.mark.parametrize("layers", [2**10 + 1, 0, 40.0]) +def test_replay_bounds_the_recorded_layer_count(layers, pending_rank, tmp_path): + rank = pending_rank + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, _ = emitted( + rank, rank._plan_flat_forward([request(65, grad=True)]), tmp_path + ) + # Replay sizes per-layer tuples from this field; refuse it before costing. + report["replay"]["memory_replay"]["rank"]["num_layers"] = layers + with pytest.raises(ValueError, match="layer count"): + reports.replay(report) + + +def test_custom_adapter_gradient_reader_is_explicitly_incomplete(monkeypatch): + from art.trainer_rank import _planner_replay + + rank = _rank(monkeypatch) + plan = rank._plan_flat_forward([_request(1)]) + original = rank._pending_adapter_gradient_bytes + rank._pending_adapter_gradient_bytes = lambda refs: original(refs) + with pytest.raises(ValueError, match="custom_runtime_estimator"): + _planner_replay.capture(rank, plan) + + @pytest.mark.parametrize( "change", ["version", "group", "layout", "gdn_segment", "budget"] ) From eedef3d110f153c30087b93670170d5b5eea753d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:16:20 +0000 Subject: [PATCH 33/40] Patch the reader with monkeypatch in the custom-reader test Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_grouped_planner_replay.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 29d384df6..ed3acc6c2 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -272,7 +272,9 @@ def test_custom_adapter_gradient_reader_is_explicitly_incomplete(monkeypatch): rank = _rank(monkeypatch) plan = rank._plan_flat_forward([_request(1)]) original = rank._pending_adapter_gradient_bytes - rank._pending_adapter_gradient_bytes = lambda refs: original(refs) + monkeypatch.setattr( + rank, "_pending_adapter_gradient_bytes", lambda refs: original(refs) + ) with pytest.raises(ValueError, match="custom_runtime_estimator"): _planner_replay.capture(rank, plan) From a4f05f82d860d9b9ca2914a2fae4d1dffab389e3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:27:30 +0000 Subject: [PATCH 34/40] Validate adapter pending facts before scanning them Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index 45c7acc22..739ee8360 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -395,7 +395,6 @@ def integer(value: Any, *, minimum: int = 0) -> None: fields(adapter, {"kind", "name", "pending"}) if ( not group["grad"] - or (adapter["kind"] is None and any(adapter["pending"] or ())) or (adapter["kind"] is not None and type(adapter["kind"]) is not str) or len(adapter["kind"] or "") > 64 or type(adapter["name"]) is not str @@ -407,6 +406,9 @@ def integer(value: Any, *, minimum: int = 0) -> None: reserve(128 + 12 * len(adapter["name"]) + 24 * len(adapter["pending"])) for value in adapter["pending"]: integer(value) + # A slot without a kind (megatron-less reference) has none pending. + if adapter["kind"] is None and any(adapter["pending"]): + raise ValueError("invalid adapter gradient facts") gdn = group["gdn"] if gdn is not None: fields(gdn, {"layers", "shapes", "segments"}) From 3fa5dbfd1093b47b3848cf3c1b2372fe5a122ab9 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:32:49 +0000 Subject: [PATCH 35/40] Refuse mixed slot kinds in facts and a zero layer count at capture Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 13 ++++++++++++- tests/unit/test_grouped_planner_replay.py | 11 +++++++++-- 2 files changed, 21 insertions(+), 3 deletions(-) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index 739ee8360..c3b7db458 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -76,7 +76,10 @@ def reserve(size: int) -> None: def capture(rank: Any, plan: Any) -> dict[str, Any]: - if rank._num_layers > MAX_LAYERS or not 0 < len(plan.groups) <= _MAX_GROUPS: + if ( + not 0 < rank._num_layers <= MAX_LAYERS + or not 0 < len(plan.groups) <= _MAX_GROUPS + ): raise ValueError("runtime_group_inventory_over_limit") if ( getattr(rank.runtime.provider, "expert_model_parallel_size", 1) > 1 @@ -436,6 +439,14 @@ def integer(value: Any, *, minimum: int = 0) -> None: ) for value in segment.values(): integer(value) + # Live slot references all have a kind, or (without megatron) none do. + kinds = { + group["adapter"]["kind"] is None + for group in groups + if group["adapter"] is not None + } + if len(kinds) > 1: + raise ValueError("invalid adapter gradient facts") if len(json.dumps(facts, separators=(",", ":"))) > _MAX_BYTES: raise ValueError("runtime_facts_over_limit") diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index ed3acc6c2..5ee494755 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -231,7 +231,8 @@ def test_selected_adapter_gradients_are_replayed_and_frozen(layer, tmp_path): @pytest.mark.parametrize( - "change", ["gradient", "length", "value", "kind_length", "kindless_pending"] + "change", + ["gradient", "length", "value", "kind_length", "kindless_pending", "mixed_kinds"], ) def test_adapter_fact_validation_rejects_forged_input(change, layer, tmp_path): report, _, _ = adapter_report(layer, tmp_path) @@ -247,8 +248,14 @@ def test_adapter_fact_validation_rejects_forged_input(change, layer, tmp_path): adapter["pending"][0] = 1.5 elif change == "kind_length": adapter["kind"] = "k" * 65 - else: + elif change == "kindless_pending": adapter["kind"] = None + else: + groups = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "groups" + ] + groups[0]["grad"] = True + groups[0]["adapter"] = {"kind": None, "name": "base", "pending": []} with pytest.raises(ValueError): reports.replay(report) From c41148ec8c0077aa1504970d643a833fbad7a07e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:44:24 +0000 Subject: [PATCH 36/40] Name the refusal each forged adapter fact must hit Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_grouped_planner_replay.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 5ee494755..c0c4b1ba8 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -256,7 +256,13 @@ def test_adapter_fact_validation_rejects_forged_input(change, layer, tmp_path): ] groups[0]["grad"] = True groups[0]["adapter"] = {"kind": None, "name": "base", "pending": []} - with pytest.raises(ValueError): + # Name the refusal: a forged fact must fail validation, not a later check. + message = ( + "invalid runtime dimension" + if change == "value" + else "invalid adapter gradient facts" + ) + with pytest.raises(ValueError, match=message): reports.replay(report) From e6edc930598d933efde767ee0ef05ab8831d60dd Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 05:22:09 +0000 Subject: [PATCH 37/40] Price the recomputed mixer and staged head backward on the extracted layout Port #963 (through 35850f744) onto ART's extracted planner modules: the recomputed layer's attention/GDN activations, routing state and TE workspaces in the checkpoint floor; routed rows per group; and the head backward priced as its own stage where traced. The planner-miss replay freezes the new model and process readers (facts version 3) and records the plan's head staging and routed rows. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/megatron/context_parallel/runtime.py | 39 ++ src/art/trainer_rank/_impl.py | 319 ++++++++-- src/art/trainer_rank/_memory.py | 556 +++++++++++++++--- src/art/trainer_rank/_micro_batch_planner.py | 166 +++++- src/art/trainer_rank/_planner_replay.py | 131 ++++- tests/unit/test_grouped_planner_replay.py | 90 +++ .../test_trainer_rank_admission_inputs.py | 1 + ...trainer_rank_checkpoint_gradient_memory.py | 184 +++++- .../test_trainer_rank_checkpoint_memory.py | 184 +++++- .../test_trainer_rank_converted_memory.py | 66 ++- tests/unit/test_trainer_rank_head_memory.py | 26 +- .../test_trainer_rank_head_stage_memory.py | 456 ++++++++++++++ tests/unit/test_trainer_rank_moe_memory.py | 125 +++- .../unit/test_trainer_rank_pending_memory.py | 20 +- tests/unit/test_trainer_rank_shared_memory.py | 35 +- tests/unit/test_trainer_rank_slot_memory.py | 32 +- 16 files changed, 2204 insertions(+), 226 deletions(-) create mode 100644 tests/unit/test_trainer_rank_head_stage_memory.py diff --git a/src/art/megatron/context_parallel/runtime.py b/src/art/megatron/context_parallel/runtime.py index 821b60e76..531cb8829 100644 --- a/src/art/megatron/context_parallel/runtime.py +++ b/src/art/megatron/context_parallel/runtime.py @@ -433,6 +433,45 @@ def context_parallel_rank_model_token_counts( ) +def context_parallel_model_token_total( + *, + group_ids: torch.Tensor, + parent_ids: torch.Tensor, + topology: ParallelTopology, + config: ContextParallelConfig, + original_seq_len: int, + build_gdn_execution_spec: bool, + gdn_planner_config: Any | None = None, +) -> int: + """Return the CP group's model rows in its larger physical layout. + + Dispatch runs at least one row on every rank; an empty rank's padding row + passes through the model too. + """ + planning_key, bundle, _group_ids_cpu, _parent_ids_cpu = ( + _get_or_build_planning_bundle( + group_ids=group_ids, + parent_ids=parent_ids, + topology=topology, + config=config, + original_seq_len=original_seq_len, + build_gdn_execution_spec=build_gdn_execution_spec, + ) + ) + total = sum( + max(1, count) for count in bundle.token_layout_index.token_counts_by_rank + ) + if not build_gdn_execution_spec: + return total + decision = _plan_gdn_global_execution( + planning_key=planning_key, + bundle=bundle, + topology=topology, + gdn_planner_config=gdn_planner_config, + ) + return max(total, sum(max(1, count) for count in decision.gdn_token_counts_by_rank)) + + def _normalized_chunk_size( *, valid_tokens: int, diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 8fcf2497a..2ec11b766 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -129,6 +129,30 @@ class TopK: # 32 MiB transients live at its peak (Qwen3.6-35B-A3B CP2: the RoPE frequencies # and a frozen linear's output, at 2k to 20k tokens); warm waves do not. _COLD_RECOMPUTE_TRANSIENT_BYTES = 64 * 2**20 + +# Per local row, index state the backward keeps beside the RoPE embedding and +# hidden-width tensors: int64 positions and row maps, CP block masks and GDN +# exchange plans. Qwen3.6-35B-A3B CP2 traces at the head's backward peak: +# 106-181 bytes per row at 3.5k-8.7k rows. +_BACKWARD_ROW_STATE_BYTES = 256 + +# A head chunk below the fused statistics' row minimum takes the FP32 fallback +# (_vocab_parallel_log_z): the BF16 logits, their FP32 copy, the shifted copy +# and its saved exponent in the recompute, then the exponent's gradient, its +# BF16 cast and the target gather's gradient in backward. At most this many +# BF16 logits-sized buffers of that chunk (about 18 bytes per logit); derived +# from the code, not traced. +_HEAD_FALLBACK_BUFFERS = 9 + +# Which fused head statistics kernels have run in this process, and whether any +# call fell back to FP32 after an error: staging trusts only a proven path. +_TRITON_STATS_STATE: dict[str, Any] = {"succeeded": set(), "failed": False} + +# Set while a gradient group of a plan priced with a staged head projects it; +# per thread, and captured into each chunk's checkpoint for its recompute. +_HEAD_STATISTICS_STRICT: ContextVar[bool] = ContextVar( + "trainer_rank_head_statistics_strict", default=False +) _PLANNER_REFINEMENT_BUDGET = 2_000 _LAYOUT_SELECTION_CACHE_LIMIT = 64 @@ -1553,8 +1577,30 @@ def _expert_lora_weight_storage( return (transposes if a.shape[2] < 8 else 0, transposes, effective) -# Routed rows per rank at EP>1, relative to balanced routing (see below). -_EP_ROUTED_ROW_ALLOWANCE = 1.5 +# Routed rows on the most loaded rank at EP>1, relative to its balanced share. +# Expert-shard load is uneven per layer, from the router's expert preferences, +# and larger batches do not average it away. Qwen3.6-35B-A3B on 3.5M tokens of +# retail agent trajectories, worst layer in 200k-token batches at EP2 / EP4 / +# EP8: pretrained up to 1.22 / 1.40 / 1.62, a trained policy up to 1.24 / 1.41 / +# 1.61; a small rollout sample reached 1.95 at EP8. One production EP2 run was +# inferred at 1.35. These samples bound what was measured, not all routing. +# Unmeasured EP sizes use the next measured one; above EP8 the allowance grows +# with log2(EP) up to EP itself (every pair on one rank). +_EP_ROUTED_ROW_ALLOWANCE = {2: 1.4, 4: 1.6, 8: 2.0} + + +def _ep_routed_row_allowance(ep: int) -> float: + if ep <= 1: + return 1.0 + for size, allowance in sorted(_EP_ROUTED_ROW_ALLOWANCE.items()): + if ep <= size: + return allowance + return min(float(ep), 2.0 + 0.4 * math.log2(ep / 8)) + + +# Transformer Engine's Hopper cuBLAS workspaces: one per grouped-GEMM stream +# (four) plus the plain GEMM's, each 32 MiB + 1 KiB. +_TE_CUBLAS_WORKSPACE_BYTES = 5 * (32 * 2**20 + 1024) def _moe_dispatcher_supported( @@ -1571,6 +1617,29 @@ def _moe_dispatcher_supported( ) +def _ep_group_is_cp_group(shape: ParallelShape) -> bool: + """Whether this rank's expert-parallel group is exactly its CP group. + + HybridEP then dispatches that CP group's rows across it: at balanced + routing each rank receives the group's rows over EP, however CP split them. + """ + if shape.ep <= 1 or shape.ep != shape.cp or (shape.tp, shape.etp) != (1, 1): + return False + if not dist.is_available() or not dist.is_initialized(): + return False + try: + from megatron.core import parallel_state as ps + except ModuleNotFoundError: + return False + expert = ps.get_expert_model_parallel_group(check_initialized=False) + context = ps.get_context_parallel_group(check_initialized=False) + if expert is None or context is None: + return False + return sorted(dist.get_process_group_ranks(expert)) == sorted( + dist.get_process_group_ranks(context) + ) + + def _hybridep_rows_per_rank(capacity: int, ranks: int) -> int: """HybridEP's allocated rows per rank: TMA-aligned, at least 512, padded to the 64-row combine chunk.""" @@ -1602,8 +1671,15 @@ def _moe_output_bytes_per_token( checkpoint_grad: bool = False, converted_stages: list[tuple[int, int]] | None = None, slot_ref: "LoRASlotRef | None" = None, + shared_bytes: list[int] | None = None, + enclosed: list[bool] | None = None, ) -> int: - """Known routed-expert working set, not a complete model/compiled bound.""" + """Known routed-expert working set, not a complete model/compiled bound. + + ``shared_bytes`` collects each layer's shared-expert part of the per-token + coefficient and stages; that part follows local rows, not routed rows. + ``enclosed`` records, per MoE layer, whether its FC1 stage is priced too. + """ # CP shards rows, not the per-token working set. At EP>1 only ART's # HybridEP flex dispatcher is modeled; TP and ETP are not. if (shape.tp, shape.etp) != (1, 1): @@ -1624,9 +1700,12 @@ def _moe_output_bytes_per_token( from art.megatron.lora import LoRA, MLPExpertsLinearFC1LoRA, MLPExpertsLinearFC2LoRA # HybridEP hands each rank the pairs routed to its local experts, already - # permuted. Balanced routing gives local tokens x top-k, as at EP1; a - # pretrained CP2/EP2 run put about 1.35x that on one rank. - routed_allowance = _EP_ROUTED_ROW_ALLOWANCE if shape.ep > 1 else 1 + # permuted. Balanced routing gives local tokens x top-k, as at EP1. + routed_allowance = _ep_routed_row_allowance(shape.ep) + # Routed H-wide inputs held at the expert stage. The EP1 all-to-all path + # keeps its permuted rows and their expert-sorted copy; HybridEP permutes + # while it dispatches and returns one tensor. + dispatched = 1 if shape.ep > 1 else 2 coefficient = 0 for chunk in model: for layer in chunk.modules(): @@ -1720,8 +1799,6 @@ def _moe_output_bytes_per_token( and fc1.fused_gate_up and not fc1.non_gated and fc1.out_features == 2 * inputs.shape[-2] - # HybridEP keeps one dispatched H-wide input where the - # EP1 all-to-all keeps two; charging two is conservative. and getattr(dispatcher, "ep_size", None) == shape.ep and getattr(dispatcher, "tp_size", None) == 1 and getattr(dispatcher, "num_local_experts", 0) > 1 @@ -1730,10 +1807,10 @@ def _moe_output_bytes_per_token( and getattr(experts, "offload_moe_act", None) is False and getattr(experts, "activation_recompute", None) is False ): - # The two dispatched H-wide inputs and FC1 gate/up sum - # remain live at the FC2 sum, including in the observed - # compiled path. This is one stage, not a backward bound. - features += 2 * fc2.out_features + fc1.out_features + # The dispatched H-wide inputs and FC1 gate/up sum remain + # live at the FC2 sum, including in the observed compiled + # path. This is one stage, not a backward bound. + features += dispatched * fc2.out_features + fc1.out_features enclosing_fc1 = fc1 shared = _shared_expert_output_bytes_per_token(layer) if ( @@ -1745,10 +1822,15 @@ def _moe_output_bytes_per_token( # Gate-score backward saves a distinct pre-gate X. Charge it # beside this layer's returned X, not another layer's maximum. shared += shared - routed_rows = math.ceil(config.moe_router_topk * routed_allowance) - row_bytes = routed_rows * features * weights.element_size() + shared + if shared_bytes is not None: + shared_bytes.append(shared) + routed_rows = config.moe_router_topk * routed_allowance + row_bytes = ( + math.ceil(routed_rows * features * weights.element_size()) + shared + ) coefficient = max(coefficient, row_bytes) storage = _expert_lora_weight_storage(lora, slot_ref) + fc1_stages = False if converted_stages is not None and storage is not None: padded, transposes, effective = storage saved_fc1, rank_fc1 = 0, 0 @@ -1777,13 +1859,14 @@ def _moe_output_bytes_per_token( and first_tensors[1].shape[2] == enclosing_fc1.out_features ): first_padding, first_transposes, first_rank = first - # FC1 retains both routed H inputs and its base O1 + fc1_stages = True + # FC1 retains the routed H inputs and its base O1 # while producing adapter O1. Its sum is not live yet. converted_stages.append( ( routed_size * ( - 2 * fc2.out_features + dispatched * fc2.out_features + 2 * enclosing_fc1.out_features + first_rank ) @@ -1797,7 +1880,7 @@ def _moe_output_bytes_per_token( ( routed_size * ( - 2 * fc2.out_features + dispatched * fc2.out_features + 3 * enclosing_fc1.out_features + (first_rank if checkpoint_grad else 0) ) @@ -1851,6 +1934,23 @@ def _moe_output_bytes_per_token( + 2 * (experts_count + 1) * 4, ) ) + if enclosed is not None: + # FC1 is covered when priced beside FC2 and its converted + # stages are priced too, unless it has no adapter or the + # selected slot has no FC1 tensors to convert. + adapter = getattr(enclosing_fc1, "lora", None) + inactive = adapter is None or ( + slot_ref is not None + and type(adapter) is LoRA + and "_slot" not in vars(adapter) + and _slot_lora_tensors(adapter, slot_ref) is None + ) + enclosed.append(enclosing_fc1 is not None and (fc1_stages or inactive)) + if converted_stages is not None: + # The EP allowance gives fractional routed rows; round each stage up. + converted_stages[:] = [ + (math.ceil(per_row), fixed) for per_row, fixed in converted_stages + ] return coefficient @@ -1976,9 +2076,15 @@ def memory_field(name: str, default: Any = None) -> Any: ) forward_stages: list[tuple[int, int]] = [] gradient_stages: list[tuple[int, int]] = [] + forward_shared: list[int] = [] + gradient_shared: list[int] = [] + gradient_enclosed: list[bool] = [] self._moe_output_bytes_per_token = ( _moe_output_bytes_per_token( - runtime.model, self._parallel_shape, converted_stages=forward_stages + runtime.model, + self._parallel_shape, + converted_stages=forward_stages, + shared_bytes=forward_shared, ) if self._moe_layers else 0 @@ -1992,6 +2098,8 @@ def memory_field(name: str, default: Any = None) -> Any: self._parallel_shape, checkpoint_grad=True, converted_stages=gradient_stages, + shared_bytes=gradient_shared, + enclosed=gradient_enclosed, ) if self._moe_layers else 0 @@ -2003,6 +2111,26 @@ def memory_field(name: str, default: Any = None) -> Any: self._moe_gradient_stages = ( tuple(gradient_stages) if self._moe_checkpoint_grad_bytes_per_token else () ) + self._moe_forward_shared_bytes = ( + max(forward_shared, default=0) if self._moe_output_bytes_per_token else 0 + ) + self._moe_gradient_shared_bytes = ( + max(gradient_shared, default=0) + if self._moe_checkpoint_grad_bytes_per_token + else 0 + ) + # Recompute is covered only if every decoder layer is a priced MoE + # layer whose FC1 stage is enclosed; otherwise dense MLP or FC1 work + # the floor does not price keeps the per-boundary gradient allowance. + self._moe_gradient_enclosed = ( + tuple(gradient_enclosed) + if self._moe_checkpoint_grad_bytes_per_token + else () + ) + self._moe_recompute_covered = len( + self._moe_gradient_enclosed + ) == self._num_layers and all(self._moe_gradient_enclosed) + self._ep_group_is_cp_group = _ep_group_is_cp_group(self._parallel_shape) selection = select_scoring( device_capability=capability, device_memory_bytes=device_memory, @@ -2912,14 +3040,16 @@ def _subforward_cost( logical_tokens: int, gdn_segments: int = 0, group_rows: tuple[tuple[int, bool], ...] = (), + group_routed_rows: tuple[int, ...] | None = None, slot_refs: tuple["LoRASlotRef | None", ...] | None = None, head_workspace_bytes: int = 0, + head_backward_traced: bool = False, checkpoint_floor: tuple[int, int] = (0, 0), retained_tokens: int | None = None, hybridep_growth_bytes: int = 0, ) -> _SubforwardCost: checkpoint_memory = self._checkpoint_memory_floor( - group_rows, slot_refs, gdn_segments + group_rows, slot_refs, gdn_segments, routed_rows=group_routed_rows ) required = self._estimate_required_memory_bytes_from_values( packed_tokens=packed_tokens, @@ -2928,6 +3058,7 @@ def _subforward_cost( logical_tokens=logical_tokens, gdn_segments=gdn_segments, group_rows=group_rows, + group_routed_rows=group_routed_rows, slot_refs=slot_refs, head_workspace_bytes=head_workspace_bytes, checkpoint_floor=checkpoint_floor, @@ -2947,11 +3078,11 @@ def _subforward_cost( checkpoint_floor[0], ), ) - # One logical BF16 input gradient per eligible full/uniform/1 boundary. - # This partial peak allowance is not evidence of simultaneous distinct - # backing stores, nor a bound for compiler saves or other backward work. - # Keep it out of forward retention, including the cold fallback above. - gradient = checkpoint_retained + # Input gradients live at the recomputed layer's peak; kept out of + # forward retention, including the cold fallback above. + gradient = self._checkpoint_input_gradient_bytes( + group_rows, slot_refs, retained=checkpoint_retained + ) gradient_slots = self._gradient_slots(group_rows, slot_refs) adapter_gradient = ( self._checkpoint_adapter_gradient_bytes( @@ -2963,24 +3094,29 @@ def _subforward_cost( checkpoint_retained = output_bytes + max( checkpoint_retained, checkpoint_floor[0] ) - checkpoint_workspace = max( - checkpoint_workspace, head_workspace_bytes, checkpoint_floor[1] - ) - if gradient and self._memory_profiles.get(signature) is None: - checkpoint_workspace += _COLD_RECOMPUTE_TRANSIENT_BYTES + decoder_workspace = max(checkpoint_workspace, checkpoint_floor[1]) + checkpoint_workspace = max(decoder_workspace, head_workspace_bytes) forward_required = required if gradient: + head_stage = self._checkpoint_head_stage_bytes( + head_workspace_bytes, gradient, group_rows, slot_refs + ) + peak = checkpoint_workspace + adapter_gradient + if head_stage is not None: + # A split adds its children's adapter gradients to the largest + # workspace, and one child's head can follow another's decoder + # backward: keep the whole head stage there. + checkpoint_workspace = max(decoder_workspace, head_stage) + # An untraced head's buffers can exceed head_workspace_bytes: + # it keeps the unstaged price as a floor. + staged = max(decoder_workspace + adapter_gradient, head_stage) + peak = staged if head_backward_traced else max(peak, staged) + if self._memory_profiles.get(signature) is None: + checkpoint_workspace += _COLD_RECOMPUTE_TRANSIENT_BYTES + peak += _COLD_RECOMPUTE_TRANSIENT_BYTES required = max( required, - int( - ( - checkpoint_retained - + checkpoint_workspace - + gradient - + adapter_gradient - ) - * _MEMORY_SAFETY_FACTOR - ), + int((checkpoint_retained + gradient + peak) * _MEMORY_SAFETY_FACTOR), ) # HybridEP buffer growth stays allocated through the forward and # backward peaks, but is not forward retention. @@ -3421,6 +3557,8 @@ def _telemetry_signature(cls, plan: _AnyForwardPlan) -> dict[str, object]: } def _execute_flat_plan(self, plan: _FlatForwardPlan) -> list[AnyForwardOutput]: + # A head priced as staged must not silently widen (_head_backward_traced). + staged = bool(getattr(plan, "_head_staged", False)) outputs = [ ForwardOutput(None, None, None, None, checkpoint, no_grad) for checkpoint, no_grad in plan.output_metadata @@ -3442,7 +3580,13 @@ def _execute_flat_plan(self, plan: _FlatForwardPlan) -> list[AnyForwardOutput]: with torch.set_grad_enabled(group.grad_enabled): with use_lora_slot(group.slot_ref): prepared = self._prepare_packed_forward(group.packed) - item_outputs = self._forward_packed(group.items, prepared) + strict = _HEAD_STATISTICS_STRICT.set( + staged and group.grad_enabled + ) + try: + item_outputs = self._forward_packed(group.items, prepared) + finally: + _HEAD_STATISTICS_STRICT.reset(strict) item_outputs = [ replace( output, @@ -4286,6 +4430,7 @@ def _project_vocab_parallel( from torch.utils.checkpoint import checkpoint model = _language_model(self.runtime.model[0]) + strict_statistics = _HEAD_STATISTICS_STRICT.get() max_top_k = max((int(item.request.top_k or 0) for item in items), default=0) need_log_z = any( item.labels is not None or item.request.top_k is not None for item in items @@ -4303,6 +4448,8 @@ def _project_vocab_parallel( output_weight=output_weight, need_log_z=need_log_z, max_top_k=max_top_k, + # Captured now: the backward recompute keeps this plan's mode. + strict_statistics=strict_statistics, use_reentrant=False, ) logit_start, logit_end = logit_bounds[chunk_index : chunk_index + 2] @@ -4383,6 +4530,7 @@ def _local_head_stats( output_weight: torch.Tensor | None, need_log_z: bool, max_top_k: int, + strict_statistics: bool = False, ) -> tuple[ torch.Tensor, torch.Tensor | None, @@ -4396,11 +4544,17 @@ def _local_head_stats( log_z: torch.Tensor | None = None local_topk: tuple[torch.Tensor, torch.Tensor] | None = None if need_log_z: - topk_stats = _try_triton_local_topk_stats(local_logits, k=max_top_k) + topk_stats = _try_triton_local_topk_stats( + local_logits, k=max_top_k, strict=strict_statistics + ) logsumexp_stats = ( cast( tuple[torch.Tensor, torch.Tensor] | None, - _try_triton_stats("local_logsumexp_stats", local_logits), + _try_triton_stats( + "local_logsumexp_stats", + local_logits, + strict=strict_statistics, + ), ) if topk_stats is None else None @@ -4741,7 +4895,56 @@ def _gather_tensor_parallel_logits(self, logits: torch.Tensor) -> torch.Tensor: _record_split_memory_floor = _memory._record_split_memory_floor _split_plan_memory_check = _memory._split_plan_memory_check _head_workspace_bytes = _memory._head_workspace_bytes + + def _cp_group_model_tokens( + self, + batch: PrefixTreePack, + *, + topology: "ParallelTopology", + ) -> int: + """The CP group's model rows in its larger physical layout.""" + from art.megatron.context_parallel.runtime import ( + context_parallel_model_token_total, + ) + from art.megatron.training.microbatches import ( + _context_parallel_config_for_provider, + _gdn_planner_config_for_provider, + ) + + handler = self.runtime.model_support_handler + return context_parallel_model_token_total( + group_ids=batch.group_ids, + parent_ids=batch.parent_ids, + topology=topology, + config=_context_parallel_config_for_provider( + self.runtime.provider, + self.device, + handler, + ), + original_seq_len=int(batch.tokens.shape[1]), + build_gdn_execution_spec=handler.build_gdn_execution_spec, + gdn_planner_config=_gdn_planner_config_for_provider( + self.runtime.provider, handler + ), + ) + _group_head_workspace_bytes = _memory._group_head_workspace_bytes + _triton_min_rows = _memory._triton_min_rows + _head_backward_traced = _memory._head_backward_traced + _te_workspace_growth_bytes = _memory._te_workspace_growth_bytes + _moe_checkpoint_state_bytes_per_token = ( + _memory._moe_checkpoint_state_bytes_per_token + ) + _moe_recompute_covered_for = _memory._moe_recompute_covered_for + _checkpoint_input_gradient_bytes = _memory._checkpoint_input_gradient_bytes + _checkpoint_gradient_covered = _memory._checkpoint_gradient_covered + _checkpoint_head_stage_bytes = _memory._checkpoint_head_stage_bytes + _backward_row_state_bytes = _memory._backward_row_state_bytes + _mixer_activation_widths = _memory._mixer_activation_widths + _recomputed_mixer_bytes_per_token = _memory._recomputed_mixer_bytes_per_token + _adapter_gradient_head = staticmethod(_memory._adapter_gradient_head) + _plan_head_backward_traced = _micro_batch_planner._plan_head_backward_traced + _plan_group_routed_rows = _micro_batch_planner._plan_group_routed_rows _plan_head_workspace_bytes = _memory._plan_head_workspace_bytes _plan_hybridep_growth_bytes = _memory._plan_hybridep_growth_bytes _checkpoint_moe_bytes_per_token = _memory._checkpoint_moe_bytes_per_token @@ -4894,6 +5097,18 @@ def _active_logical_tokens(requests: Sequence[AnyForwardInput]) -> int: _PACKED_PRICED_MIN_REQUEST_TOKENS = 64 +def _traced_states(traced: Sequence[bool | None]) -> tuple[bool, ...]: + """Head stagings a plan's gradient groups allow (``_head_backward_traced``). + + Staged only when every group is traced; both when any is undecided. + """ + if not traced or any(state is False for state in traced): + return (False,) + if all(state is True for state in traced): + return (True,) + return (True, False) + + def _packed_priced(signature: "_MemorySignature", one_layer_recompute: bool) -> bool: # Measured only under one-layer full recompute: other recompute modes keep # more GDN states and activations live per segment and logical row. @@ -5428,6 +5643,7 @@ def _try_triton_local_topk_stats( local_logits: torch.Tensor, *, k: int, + strict: bool = False, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] | None: if k <= 0 or k > int( os.environ.get("ART_TRAINER_RANK_TRITON_FUSED_TOPK_MAX", "10") @@ -5438,6 +5654,7 @@ def _try_triton_local_topk_stats( _try_triton_stats( "local_topk_stats", local_logits, + strict=strict, k=min(k, int(local_logits.shape[1])), ), ) @@ -5446,8 +5663,16 @@ def _try_triton_local_topk_stats( def _try_triton_stats( name: str, local_logits: torch.Tensor, + *, + strict: bool = False, **kwargs: object, ) -> object | None: + """The fused statistics, or None for the FP32 fallback. + + ``strict``: a plan whose price relied on them raises instead of falling + back after an error (``_head_backward_traced``). Too few rows still fall + back: the head stage prices that (``_HEAD_FALLBACK_BUFFERS``). + """ if not local_logits.is_cuda: return None if os.environ.get("ART_TRAINER_RANK_TRITON_TOPK", "1").lower() in { @@ -5460,11 +5685,19 @@ def _try_triton_stats( try: from art.trainer_rank import topk - return getattr(topk, name)(local_logits, **kwargs) - except Exception: + result = getattr(topk, name)(local_logits, **kwargs) + except Exception as error: + _TRITON_STATS_STATE["failed"] = True + if strict: + raise RuntimeError( + "Fused head statistics failed in a plan admitted on their " + "memory; the FP32 fallback would exceed its price" + ) from error if os.environ.get("ART_TRAINER_RANK_TRITON_TOPK", "1").lower() == "strict": raise return None + _TRITON_STATS_STATE["succeeded"].add(name) + return result def _vocab_parallel_topk_from_local( diff --git a/src/art/trainer_rank/_memory.py b/src/art/trainer_rank/_memory.py index 40ea5b5b7..6a4654222 100644 --- a/src/art/trainer_rank/_memory.py +++ b/src/art/trainer_rank/_memory.py @@ -438,13 +438,14 @@ def _moe_workspace_terms( *, checkpoint_grad: bool = False, slot_ref: "LoRASlotRef | None" = None, -) -> tuple[int, tuple[tuple[int, int], ...]]: +) -> tuple[int, tuple[tuple[int, int], ...], int]: """Maximum of same-layer affine stages, not a retained multi-layer bank. The constructor cache covers original tensors. Explicit slots are repriced from their tensor metadata and original owners, including this rank's exact dispatcher wrapper. Ordinary non-checkpoint gradients - retain only forward-stage coverage. + retain only forward-stage coverage. The third term is the shared + expert's part of the coefficient, which stays on the local rows. """ coefficient = ( self._checkpoint_moe_bytes_per_token() @@ -456,8 +457,16 @@ def _moe_workspace_terms( "_moe_gradient_stages" if checkpoint_grad else "_moe_forward_stages", (), ) + shared = getattr( + self, + "_moe_gradient_shared_bytes" + if checkpoint_grad + else "_moe_forward_shared_bytes", + 0, + ) if slot_ref is not None and slot_ref.name is not None: selected: list[tuple[int, int]] = [] + slot_shared: list[int] = [] coefficient = ( _impl._moe_output_bytes_per_token( self.runtime.model, @@ -465,11 +474,13 @@ def _moe_workspace_terms( checkpoint_grad=checkpoint_grad, converted_stages=selected, slot_ref=slot_ref, + shared_bytes=slot_shared, ) if self._moe_layers else 0 ) stages = tuple(selected) if coefficient else () + shared = max(slot_shared, default=0) if coefficient else 0 if type(stages) is not tuple or any( type(stage) is not tuple or len(stage) != 2 @@ -477,20 +488,27 @@ def _moe_workspace_terms( for stage in stages ): raise ValueError("Invalid constructor converted-weight stages") - return coefficient, stages + if type(shared) is not int or not 0 <= shared <= coefficient: + raise ValueError("Invalid constructor shared-expert coefficient") + return coefficient, stages, shared def _moe_workspace_from_terms( - rows: int, terms: tuple[int, tuple[tuple[int, int], ...]] + rows: int, + terms: tuple[int, tuple[tuple[int, int], ...], int], + routed_rows: int | None = None, ) -> int: - coefficient, stages = terms - return ( + """``routed_rows`` is what this rank's experts receive at balanced routing + (``rows`` by default); the shared expert's part stays on the local rows.""" + coefficient, stages, shared = terms + routed = rows if routed_rows is None else max(0, min(rows, routed_rows)) + return (rows - routed) * shared + ( max( - rows * coefficient, - *(rows * per_row + fixed for per_row, fixed in stages), + routed * coefficient, + *(routed * per_row + fixed for per_row, fixed in stages), ) if stages and rows > 0 - else rows * coefficient + else routed * coefficient ) @@ -498,12 +516,14 @@ def _moe_workspace_bytes( self: TrainerRank, rows: int, *, + routed_rows: int | None = None, checkpoint_grad: bool = False, slot_ref: "LoRASlotRef | None" = None, ) -> int: return _moe_workspace_from_terms( rows, _moe_workspace_terms(self, checkpoint_grad=checkpoint_grad, slot_ref=slot_ref), + routed_rows, ) @@ -587,26 +607,41 @@ def _checkpoint_memory_floor( group_rows: tuple[tuple[int, bool], ...], slot_refs: tuple["LoRASlotRef | None", ...] | None = None, gdn_segments: int = 0, + routed_rows: tuple[int, ...] | None = None, ) -> tuple[int, int]: + """Conservative saved-boundary charge and one recomputed layer's workspace. + + ``routed_rows`` are each group's balanced dispatched rows per rank + (``_plan_group_routed_rows``); by default, its local rows. Eligibility is + ``_checkpoint_layers``; the arithmetic, shared with grouped CPU replay, is + ``_checkpoint_floor_from_facts``. + """ layers = _checkpoint_layers(self, group_rows) if not layers: return 0, 0 retained, workspace = _checkpoint_floor_from_facts( - self, group_rows, slot_refs, gdn_segments, layers + self, group_rows, slot_refs, gdn_segments, layers, routed_rows ) - # HybridEP runtime state is intentionally outside grouped CPU replay v1. - hybrid_rows = None + # HybridEP runtime state is intentionally outside grouped CPU replay. if ( any(grad for _, grad in group_rows) and self._topology_key() == (1, 1, 2, 1) and self._parallel_shape == _impl.ParallelShape(tp=1, cp=2, ep=2, etp=1) and self._moe_memory_supported ): - hybrid_rows = max(rows for rows, _ in group_rows) + # Recompute runs after _execute_flat_plan restores the communication + # high-water. Combine allocates a fresh BF16 [P, H] before cropping; + # this is separate from already-held native buffer capacity. Do not + # prune graph references or reset execution state while estimating. + rows = max(rows for rows, _ in group_rows) if any(ref() is not None for ref in self._pending_hybridep_graphs): - hybrid_rows = max(hybrid_rows, self._hybridep_rows_high_water) - if hybrid_rows is not None: - workspace = max(workspace, -(-hybrid_rows // 4) * 4 * self._hidden_size * 2) + rows = max(rows, self._hybridep_rows_high_water) + # The combine output and the TE workspaces are live together. + workspace = max( + workspace, + -(-rows // 4) * 4 * self._hidden_size * 2 + + self._te_workspace_growth_bytes(), + ) return retained, workspace @@ -616,32 +651,76 @@ def _checkpoint_floor_from_facts( slot_refs: tuple["LoRASlotRef | None", ...] | None, gdn_segments: int, layers: int, + routed_rows: tuple[int, ...] | None = None, ) -> tuple[int, int]: + """Count actual local full/uniform/1 boundaries, including aliases, rather + than claiming measured distinct storage. Only this call's new groups + enter the term; already-live graphs remain in the availability baseline. + The workspace is one MoE stage plus what else is live beside it. + Gradient groups recompute the layer, so its attention or GDN mixer + keeps its saved activations across the MoE stage. No-grad groups keep + decoder input, current layer input, its MLP residual and norm output. + Count these four row tensors separately from returned outputs, allowing + storage aliases. This is not a bound for custom preprocessing or all of + backward. With sequence parallelism a rank saves only its shard of each + boundary; that is priced only where ``_sequence_parallel_floor_covered`` + holds, and there, for gradient waves, the recomputed GDN layer's + recurrent states for ``gdn_segments`` (gradient groups' segments) plus + padding, as traced (the recomputed mixer is not added at TP > 1). + """ if not layers: return 0, 0 gradient_rows = sum(rows for rows, grad in group_rows if grad) _, tp, _, _ = self._topology_key() - # Physical rows are padded to a multiple of TP; each rank saves its shard. - retained = ( - sum(-(-rows // tp) for rows, grad in group_rows if grad) - * layers - * self._hidden_size - * 2 - ) - if gradient_rows: - self._checkpoint_moe_bytes_per_token() refs = (None,) * len(group_rows) if slot_refs is None else slot_refs + if tp > 1: + # The traced TP x SP floor. Physical rows are padded to a multiple of + # TP; each rank saves its shard. + retained = ( + sum(-(-rows // tp) for rows, grad in group_rows if grad) + * layers + * self._hidden_size + * 2 + ) + if gradient_rows: + self._checkpoint_moe_bytes_per_token() + workspace = max( + self._moe_workspace_bytes(rows, checkpoint_grad=grad, slot_ref=ref) + + (0 if grad else 4 * rows * self._hidden_size * 2) + for (rows, grad), ref in zip(group_rows, refs, strict=True) + ) + if self._gdn_layers and gradient_rows: + # Recurrent states grow with segments, not rows; backward recomputes + # one layer at a time. Padding to TP adds up to TP - 1 one-token + # roots per group. Kernel-internal chunk states are not bounded here. + roots = gdn_segments + (tp - 1) * sum(grad for _, grad in group_rows) + workspace += math.ceil(roots * self._gdn_segment_layer_bytes()) + return retained, workspace + # Boundaries on the busiest rank's rows and one recomputed layer. + routed = (None,) * len(group_rows) if routed_rows is None else routed_rows + retained = gradient_rows * layers * self._hidden_size * 2 + moe = self._checkpoint_moe_bytes_per_token() if gradient_rows else 0 + # Beside the mixer, the recomputed layer keeps its post-mixer residual + # and pre-MLP norm output, and its MoE stage its routing state. + mixer = ( + self._recomputed_mixer_bytes_per_token() + + ( + 2 * self._hidden_size * 2 + self._moe_checkpoint_state_bytes_per_token() + if moe + else 0 + ) + if gradient_rows + else 0 + ) workspace = max( - self._moe_workspace_bytes(rows, checkpoint_grad=grad, slot_ref=ref) - + (0 if grad else 4 * rows * self._hidden_size * 2) - for (rows, grad), ref in zip(group_rows, refs, strict=True) + self._moe_workspace_bytes( + rows, routed_rows=dispatched, checkpoint_grad=grad, slot_ref=ref + ) + + (mixer * rows if grad else 4 * rows * self._hidden_size * 2) + for (rows, grad), ref, dispatched in zip(group_rows, refs, routed, strict=True) ) - if tp > 1 and self._gdn_layers and gradient_rows: - # Recurrent states grow with segments, not rows; backward recomputes - # one layer at a time. Padding to TP adds up to TP - 1 one-token - # roots per group. Kernel-internal chunk states are not bounded here. - roots = gdn_segments + (tp - 1) * sum(grad for _, grad in group_rows) - workspace += math.ceil(roots * self._gdn_segment_layer_bytes()) + if moe: + workspace += self._te_workspace_growth_bytes() return retained, workspace @@ -765,13 +844,17 @@ def _checkpoint_gradient_groups( def _checkpoint_adapter_gradient_bytes( - self: TrainerRank, groups: Sequence[tuple[LoRASlotRef | None, Sequence[int]]] + self: TrainerRank, + groups: Sequence[tuple[LoRASlotRef | None, Sequence[int]]], + *, + head: bool = False, ) -> int: """The recompute backward's adapter-gradient peak beyond released boundaries. ``groups`` gives each gradient group's adapter slot (None for the base model) and each decoder layer's saved-boundary bytes - (``_adapter_gradient_walk``). + (``_adapter_gradient_walk``). With ``head``, those live while a group's + head runs its backward instead (``_adapter_gradient_head``). """ chains = [] for slot, boundaries in groups: @@ -779,6 +862,8 @@ def _checkpoint_adapter_gradient_bytes( if pending and len(pending) != len(boundaries) + 1: return 0 chains.append((pending or (0,) * (len(boundaries) + 1), boundaries)) + if head: + return self._adapter_gradient_head(chains) return self._adapter_gradient_walk(chains) @@ -821,6 +906,336 @@ def _adapter_gradient_walk( return worst +def _triton_min_rows(self: TrainerRank) -> int: + """Rows the fused head statistics need in a chunk (process-wide setting).""" + return int(_impl.os.environ.get("ART_TRAINER_RANK_TRITON_MIN_ROWS", "64")) + + +def _head_backward_traced( + self: TrainerRank, + requests: Sequence[AnyForwardInput], + rows: int, + upper_rows: int | None = None, +) -> bool | None: + """Whether a gradient group's head backward is the traced one. + + The head stage (``_checkpoint_head_stage_bytes``) relies on + ``_group_head_workspace_bytes`` bounding the head's backward buffers. + Qwen3.6-35B-A3B CP2 traces establish that for target-only requests on + the standard logit scale through the fused Triton statistics, which + need at least ``ART_TRAINER_RANK_TRITON_MIN_ROWS`` rows in the first + projected chunk. Top-k, logits and hidden-state outputs keep further + dense gradients, and the FP32 fallback wider copies; other CP sizes + are untraced. The fused statistics must have run in this process and + never fallen back after an error; a plan priced on them then runs them + strictly (``_plan_head_backward_traced``), so a later failure raises + rather than taking the wider FP32 path. ``rows`` bounds the first chunk's rows from below, + ``upper_rows`` from above: None when they straddle the threshold. + A CP rank projecting fewer rows falls back on its own; the head stage + prices that (``_HEAD_FALLBACK_BUFFERS``). ``ART_TRAINER_RANK_TRITON_*`` + settings are process-wide: changing them while a plan is in flight is + unsupported (a staged chunk could then fall back unpriced). + """ + if ( + not requests + or any( + request.target_tokens is None + # Several labels per row save gather indices and masks per label; + # one label per input token, in either accepted layout, does not. + or request.target_tokens.numel() != request.input_tokens.numel() + or request.top_k is not None + or request.logits + or request.hidden_states + for request in requests + ) + or _impl.os.environ.get("ART_TRAINER_RANK_TRITON_TOPK", "1").lower() + in {"0", "false"} + # Target-only heads run the logsumexp kernel; another kernel's + # success does not prove it. + or "local_logsumexp_stats" not in _impl._TRITON_STATS_STATE["succeeded"] + or _impl._TRITON_STATS_STATE["failed"] + or self._topology_key()[2] != 2 + or not self._head_workspace_bytes(1) + or not _head_target_backward(self) + ): + return False + minimum = self._triton_min_rows() + if rows >= minimum: + return True + if upper_rows is None or upper_rows < minimum: + return False + return None + + +def _te_workspace_growth_bytes(self: TrainerRank) -> int: + """Transformer Engine's cuBLAS workspaces, until its GEMMs allocate them. + + The first plain and grouped GEMMs allocate them during a call, and TE + keeps them for the process; later calls see them as used memory. + """ + try: + from transformer_engine.pytorch.cpp_extensions import gemm + except ImportError: + return _impl._TE_CUBLAS_WORKSPACE_BYTES + info = getattr(gemm.get_cublas_workspace, "cache_info", None) + if callable(info) and info().currsize >= 2: + return 0 + return _impl._TE_CUBLAS_WORKSPACE_BYTES + + +def _moe_checkpoint_state_bytes_per_token(self: TrainerRank) -> int: + """Per local token beside the recomputed MoE stage's routed rows. + + FP32 router scores and the boolean routing map; the dispatcher's state + (the EP1 permutation's int32 row-id map of 2E + 1, or HybridEP's FP32 + probability copy and handle metadata of about 5E); and the shared + expert's saved FC1 gate/up and GLU outputs. Qwen3.6-35B-A3B traces: + 3,332 and 3,593 bytes of routing state at EP1 and EP2, 3 KB shared. + """ + geometry = self._geometry + experts = geometry.moe_experts + if not experts: + return 0 + routing = experts * (4 + 1) + ( + 4 * (2 * experts + 1) if self._parallel_shape.ep == 1 else 9 * experts + 16 + ) + return routing + 3 * geometry.moe_shared_expert_ffn * self._param_dtype_size + + +def _moe_recompute_covered_for( + self: TrainerRank, slot_ref: "LoRASlotRef | None" +) -> bool: + """Whether an explicit slot keeps the constructor's full MoE coverage.""" + if not getattr(self, "_moe_recompute_covered", False): + return False + if slot_ref is None or slot_ref.name is None: + return True + enclosed: list[bool] = [] + coefficient = _impl._moe_output_bytes_per_token( + self.runtime.model, + self._parallel_shape, + checkpoint_grad=True, + converted_stages=[], + slot_ref=slot_ref, + enclosed=enclosed, + ) + # The constructor's walk already matched every decoder layer. + return ( + coefficient > 0 + and len(enclosed) == len(getattr(self, "_moe_gradient_enclosed", ())) + and all(enclosed) + ) + + +def _checkpoint_input_gradient_bytes( + self: TrainerRank, + group_rows: tuple[tuple[int, bool], ...], + slot_refs: tuple["LoRASlotRef | None", ...] | None = None, + *, + retained: int | None = None, +) -> int: + """Gradient rows live at the recomputed layer's peak. + + Backward recomputes the last layer first, so its peak meets every saved + boundary but only the one incoming gradient. Where the MoE stage covers + every layer's recompute, FC1 included (Qwen3.6-35B-A3B traces at CP1, + CP2/EP1 and EP2/CP2), charge that gradient. Elsewhere, including above + CP2 (more remote attention stages than the mixer's CP2 allowance), keep + one gradient per boundary: that allowance also covers dense MLP and + other recompute work the floor does not price. ``retained`` is the + floor's boundary charge where the caller already has it. + """ + if retained is None: + retained, _ = self._checkpoint_memory_floor(group_rows) + if not retained: + return 0 + if self._checkpoint_gradient_covered(group_rows, slot_refs): + return sum(rows for rows, grad in group_rows if grad) * self._hidden_size * 2 + return retained + + +def _checkpoint_gradient_covered( + self: TrainerRank, + group_rows: tuple[tuple[int, bool], ...], + slot_refs: tuple["LoRASlotRef | None", ...] | None, +) -> bool: + """Whether the traced MoE stage covers every gradient group's recompute.""" + refs = (None,) * len(group_rows) if slot_refs is None else slot_refs + return bool( + self._topology_key()[2] <= 2 + and self._checkpoint_moe_bytes_per_token() + and all( + self._moe_recompute_covered_for(ref) + for (_, grad), ref in zip(group_rows, refs, strict=True) + if grad + ) + ) + + +def _checkpoint_head_stage_bytes( + self: TrainerRank, + head_workspace_bytes: int, + gradient: int, + group_rows: tuple[tuple[int, bool], ...], + slot_refs: tuple["LoRASlotRef | None", ...] | None, +) -> int | None: + """The head backward's peak beyond the boundaries and ``gradient``. + + The head finishes its backward before the decoder's recompute starts, + so its buffers never meet a recomputed layer's workspace or the adapter + gradients that recompute allocates. Where the MoE stage covers the + recompute, ``gradient`` is the gradient groups' rows x 2H, and each of + these is at most that: the final decoder outputs, the hidden rows each + checkpointed head chunk saved (this rank's rows), and the hidden-row + gradient. Beside them are each row's RoPE embedding and index state, + the head's own buffers, TE's cuBLAS workspaces from the forward's first + GEMMs, and the adapter gradients of groups whose backward ran first. + Qwen3.6-35B-A3B CP2 traces (EP1 and EP2, single and multi-request + waves) show these terms at the head's peak. Elsewhere None: the head + shares the decoder stage. That bounds heads whose backward is the + traced one (``_head_backward_traced``); callers keep the unstaged price + beside any other, whose buffers can exceed ``head_workspace_bytes``. + """ + if not head_workspace_bytes or not self._checkpoint_gradient_covered( + group_rows, slot_refs + ): + return None + rows = sum(rows for rows, grad in group_rows if grad) + # A rank's chunk below the fused minimum (CP splits the projected rows + # unevenly) takes the FP32 fallback. + fallback = _impl._HEAD_FALLBACK_BUFFERS * self._head_workspace_bytes( + min( + self._triton_min_rows() - 1, + _impl._HEAD_CHUNK_TOKENS, + ) + ) + return ( + max(head_workspace_bytes, fallback) + + 2 * gradient + + rows * self._backward_row_state_bytes() + + self._te_workspace_growth_bytes() + + self._checkpoint_adapter_gradient_bytes( + self._checkpoint_gradient_groups(group_rows, slot_refs), head=True + ) + ) + + +def _backward_row_state_bytes(self: TrainerRank) -> int: + """Per local row, what the backward keeps beside hidden-width tensors. + + The FP32 RoPE embedding at the rotary width (256 bytes for + Qwen3.6-35B-A3B) and ``_BACKWARD_ROW_STATE_BYTES`` of index state. + """ + rope = getattr(_impl._language_model(self.runtime.model[0]), "rotary_pos_emb", None) + frequencies = getattr(rope, "inv_freq", None) + width = ( + 2 * frequencies.numel() if isinstance(frequencies, _impl.torch.Tensor) else 0 + ) + return width * 4 + _impl._BACKWARD_ROW_STATE_BYTES + + +def _mixer_activation_widths(self: TrainerRank) -> tuple[float, float]: + """Saved activations per token of one attention and one GDN layer. + + Elements, not bytes: what a layer's attention or GDN mixer keeps for + its backward, including the input norm output (and gathered LoRA + inputs under sequence parallelism). + """ + geometry = self._geometry + hidden = self._hidden_size + tp = max(1, self._topology_key()[1]) + sp = tp if self._sequence_parallel else 1 + # Gathered LoRA inputs alias norm output without sequence sharding. + gathered = hidden if sp > 1 else 0 + common = 2 * hidden / sp + gathered + attention_width = geometry.num_attention_heads * geometry.kv_channels or hidden + kv_width = geometry.num_query_groups * geometry.kv_channels or hidden + gated = self._attention_output_gate + attention = common + ((7 if gated else 5) * attention_width + 3 * kv_width) / tp + if 0 < geometry.num_query_groups < tp: + # SelfAttentionLinearQKVLoRA constructs global QKV before + # slicing it when KV groups cannot be partitioned across TP. + attention += ((2 if gated else 1) * attention_width + 2 * kv_width) * ( + 1 - 1 / tp + ) + gdn = ( + common + + ( + 4 * geometry.gdn_key_heads * geometry.gdn_key_head_dim + + 8 * geometry.gdn_value_heads * geometry.gdn_value_head_dim + ) + / tp + ) + return attention, gdn + + +def _recomputed_mixer_bytes_per_token(self: TrainerRank) -> int: + """Saved mixer activations of the layer being recomputed, per row. + + Full recompute replays one layer with gradients, and its attention or + GDN mixer keeps what its backward needs across that layer's MoE stage. + Price the larger mixer the model has. Sites and sizes come from + Qwen3.6-35B-A3B allocator traces on H200, per local token at the + layer's recompute peak: + + - attention: the retained attention width above (66 KB measured at + CP1, 69 KB priced). A context-parallel rank also keeps its + stage-padded Q/K/V, the stage output and a core-attention copy + (94 KB measured at CP2, 95 KB priced). CP above 2 uses the CP2 + allowance; ranks with several remote stages may keep more. + - GDN: norm output, the projected q/k/v, their l2norm outputs + expanded to the value heads, five more value-width tensors (z, two + segment-layout tensors, the gated-norm output and its gated + product) and the chunk decay matrix: 80 KB measured at CP1, 82 KB + priced. A context-parallel rank's exchanged layout holds about 7% + more rows; its hidden-width input exchange and value-width output + allowance price that (88 KB measured at CP2, 94 KB priced). + """ + geometry = self._geometry + hidden = self._hidden_size + tp = max(1, self._topology_key()[1]) + cp = self._topology_key()[2] > 1 + widths = [] + if self._gdn_layers < self._num_layers: + attention, _gdn = self._mixer_activation_widths() + if cp: + q = geometry.num_attention_heads * geometry.kv_channels or hidden + kv = geometry.num_query_groups * geometry.kv_channels or hidden + attention += (3 * q + 2 * kv) / tp + widths.append(attention) + if self._gdn_layers: + key = geometry.gdn_key_heads * geometry.gdn_key_head_dim + value = geometry.gdn_value_heads * geometry.gdn_value_head_dim + normalized = 2 * geometry.gdn_value_heads * geometry.gdn_key_head_dim + chunk = 64 * geometry.gdn_value_heads + gdn = hidden + (2 * key + normalized + 6 * value + chunk) / tp + if cp: + gdn += hidden + value / tp + widths.append(gdn) + return int(max(widths, default=0) * self._param_dtype_size) + + +def _adapter_gradient_head( + chains: Sequence[tuple[Sequence[int], Sequence[int]]], +) -> int: + """Adapter gradients live while one group's head runs its backward. + + Chains as ``_adapter_gradient_walk``. None of the group's decoder layers + has run yet; its gradients outside the decoder are live, and so is every + other group whose gradients outweigh its boundaries, as any may have run + first. + """ + nets = [max(0, sum(pending) - sum(boundaries)) for pending, boundaries in chains] + others = sum(nets) + return max( + ( + pending[-1] + others - net + for (pending, _), net in zip(chains, nets, strict=True) + ), + default=0, + ) + + def _retained_memory_bytes( self: TrainerRank, signature: _impl._MemorySignature, @@ -870,6 +1285,7 @@ def _estimate_flat_forward( memory_minimal: bool = False, sync_planning_errors: bool = False, gdn_segments: list[int] | None = None, + head_traced: list[bool | None] | None = None, ) -> tuple[int, int, _impl._MemorySignature, tuple[tuple[int, bool], ...], int] | None: """Estimate packed tokens for width probing. @@ -885,6 +1301,8 @@ def _estimate_flat_forward( ``gdn_segments`` receives each gradient group's segment count: exact layouts' actual counts; in cheap mode, the same kind of bound as the token count (twice the requests, as a radix tree has fewer, or one). + ``head_traced`` receives each gradient group's ``_head_backward_traced`` + over its projected-row bounds. """ if sync_planning_errors: @@ -933,6 +1351,10 @@ def _estimate_flat_forward( head_requests = tuple(requests[index] for index in group_indices) lower = self._head_projection_rows(head_requests, lower_bound=True) upper = self._head_projection_rows(head_requests) + if grad_enabled and head_traced is not None: + head_traced.append( + self._head_backward_traced(head_requests, lower, upper) + ) if exact: tree, layout = self._select_group_layout( tuple( @@ -1161,8 +1583,10 @@ def _memory_check( logical_tokens=forward.active_logical_tokens, gdn_segments=forward.grad_segment_count, group_rows=self._plan_group_rows(forward), + group_routed_rows=self._plan_group_routed_rows(forward), slot_refs=tuple(g.slot_ref for g in forward.groups), head_workspace_bytes=self._plan_head_workspace_bytes(forward), + head_backward_traced=self._plan_head_backward_traced(forward), checkpoint_floor=_impl._gdn_memory.plan_floor(self, forward), retained_tokens=self._plan_retained_tokens(forward), ) + int(self._plan_hybridep_growth_bytes(forward) * _impl._MEMORY_SAFETY_FACTOR) @@ -1273,8 +1697,10 @@ def _estimate_required_memory_bytes_from_values( logical_tokens: int | None = None, gdn_segments: int = 0, group_rows: tuple[tuple[int, bool], ...] = (), + group_routed_rows: tuple[int, ...] | None = None, slot_refs: tuple["LoRASlotRef | None", ...] | None = None, head_workspace_bytes: int = 0, + head_backward_traced: bool = False, checkpoint_floor: tuple[int, int] = (0, 0), retained_tokens: int | None = None, include_checkpoint_input_gradient: bool = True, @@ -1295,24 +1721,7 @@ def _estimate_required_memory_bytes_from_values( # Gathered LoRA inputs alias norm output without sequence sharding. gathered = hidden if sp > 1 else 0 common = 2 * hidden / sp + gathered - attention_width = geometry.num_attention_heads * geometry.kv_channels or hidden - kv_width = geometry.num_query_groups * geometry.kv_channels or hidden - gated = self._attention_output_gate - attention = common + ((7 if gated else 5) * attention_width + 3 * kv_width) / tp - if 0 < geometry.num_query_groups < tp: - # SelfAttentionLinearQKVLoRA constructs global QKV before - # slicing it when KV groups cannot be partitioned across TP. - attention += ((2 if gated else 1) * attention_width + 2 * kv_width) * ( - 1 - 1 / tp - ) - gdn = ( - common - + ( - 4 * geometry.gdn_key_heads * geometry.gdn_key_head_dim - + 8 * geometry.gdn_value_heads * geometry.gdn_value_head_dim - ) - / tp - ) + attention, gdn = self._mixer_activation_widths() ffn_width = geometry.ffn_hidden_size or 4 * hidden mlp = common + self._mlp_activation_factor * ffn_width / tp if geometry.moe_experts: @@ -1361,37 +1770,48 @@ def _estimate_required_memory_bytes_from_values( ) # Groups execute sequentially: summed packed rows conservatively bound # this FC2 component, not all workspace or retained graphs. + grouped = signature.topology[2] > 1 and bool(group_rows) static_compute = max( static_compute, *( self._moe_workspace_bytes( - sum(rows for rows, _ in group_rows) - if signature.topology[2] > 1 and group_rows - else packed_tokens, + sum(rows for rows, _ in group_rows) if grouped else packed_tokens, + routed_rows=sum(group_routed_rows) + if grouped and group_routed_rows is not None + else None, slot_ref=ref, ) for ref in (slot_refs or (None,)) ), ) retained, workspace = ( - self._checkpoint_memory_floor(group_rows, slot_refs, gdn_segments) + self._checkpoint_memory_floor( + group_rows, slot_refs, gdn_segments, routed_rows=group_routed_rows + ) if checkpoint_memory is None else checkpoint_memory ) - backward = 0 + decoder_workspace = max(workspace, checkpoint_floor[1]) + peak = max(decoder_workspace, head_workspace_bytes) if include_checkpoint_input_gradient and retained: - # The backward's other end and cold transients, as _subforward_cost. - backward = retained + self._checkpoint_adapter_gradient_bytes( + # The input gradient, the backward's other end and cold transients, + # staged as _subforward_cost. + gradient = self._checkpoint_input_gradient_bytes(group_rows, slot_refs) + adapter_gradient = self._checkpoint_adapter_gradient_bytes( self._checkpoint_gradient_groups(group_rows, slot_refs) ) + head_stage = self._checkpoint_head_stage_bytes( + head_workspace_bytes, gradient, group_rows, slot_refs + ) + peak += adapter_gradient + if head_stage is not None: + # An untraced head keeps the unstaged price as a floor. + staged = max(decoder_workspace + adapter_gradient, head_stage) + peak = staged if head_backward_traced else max(peak, staged) + peak += gradient if profiled is None and any(grad for _, grad in group_rows): - backward += _impl._COLD_RECOMPUTE_TRANSIENT_BYTES - static_compute = max( - static_compute, - max(retained, checkpoint_floor[0]) - + max(workspace, head_workspace_bytes, checkpoint_floor[1]) - + backward, - ) + peak += _impl._COLD_RECOMPUTE_TRANSIENT_BYTES + static_compute = max(static_compute, max(retained, checkpoint_floor[0]) + peak) if signature.topology[2] > 1: # Local head results coexist with full CP outputs during gathering. # Uneven rank plans can assign all of an item's rows to one rank. diff --git a/src/art/trainer_rank/_micro_batch_planner.py b/src/art/trainer_rank/_micro_batch_planner.py index 2a077b57d..5e2b7bfa8 100644 --- a/src/art/trainer_rank/_micro_batch_planner.py +++ b/src/art/trainer_rank/_micro_batch_planner.py @@ -431,6 +431,7 @@ def _split_chunk_lower_cost( packed_tokens = 0 unshared_packed_tokens = 0 head_workspace_bytes = 0 + head_traced: list[bool | None] = [] group_rows: list[tuple[int, bool]] = [] for (_slot, grad_enabled), group_indices in groups: estimated = _impl.estimate_prefix_tree_packed_tokens( @@ -444,15 +445,24 @@ def _split_chunk_lower_cost( cp = max(1, self._topology_key()[2]) group_rows.append((-(-physical_rows // cp), grad_enabled)) head_requests = tuple(requests[index] for index in group_indices) + lower = self._head_projection_rows(head_requests, lower_bound=True) head_workspace_bytes = max( head_workspace_bytes, self._group_head_workspace_bytes( - self._head_projection_rows(head_requests, lower_bound=True), + lower, head_requests, grad_enabled=grad_enabled, lower_bound=True, ), ) + if grad_enabled: + head_traced.append( + self._head_backward_traced( + head_requests, + lower, + self._head_projection_rows(head_requests), + ) + ) unshared_packed_tokens += self._physical_tokens( sum(int(rows[index].numel()) for index in group_indices) ) @@ -464,17 +474,27 @@ def _split_chunk_lower_cost( slot_groups=tuple(key for key, _ in groups), ) logical_tokens = _impl._active_logical_tokens(requests) - cost = self._subforward_cost( - packed_tokens=packed_tokens, - output_bytes=output_bytes, - signature=signature, - logical_tokens=logical_tokens, - group_rows=tuple(group_rows), - slot_refs=tuple(ref for (ref, _), _ in groups), - head_workspace_bytes=head_workspace_bytes, - # The average CP load is an optimistic bound, not an admission cost. - retained_tokens=(packed_tokens + signature.topology[2] - 1) - // signature.topology[2], + # Whether the exact plan stages its head can depend on projected rows + # these bounds leave open; the lower of both prices bounds either. + cost = min( + ( + self._subforward_cost( + packed_tokens=packed_tokens, + output_bytes=output_bytes, + signature=signature, + logical_tokens=logical_tokens, + group_rows=tuple(group_rows), + slot_refs=tuple(ref for (ref, _), _ in groups), + head_workspace_bytes=head_workspace_bytes, + head_backward_traced=traced, + # The average CP load is an optimistic bound, not an + # admission cost. + retained_tokens=(packed_tokens + signature.topology[2] - 1) + // signature.topology[2], + ) + for traced in _impl._traced_states(head_traced) + ), + key=lambda cost: cost.required, ) profile = self._memory_profiles.get(signature) if ( @@ -535,6 +555,76 @@ def _plan_group_rows( ) +def _plan_head_backward_traced(self: TrainerRank, plan: _FlatForwardPlan) -> bool: + """Every gradient group's head backward is traced (``_head_backward_traced``). + + When that stages the plan's head price (``_checkpoint_head_stage_bytes`` + applies), marks the plan: its gradient groups then run the fused + statistics strictly (``_execute_flat_plan``). The mark stays once set: + a later re-pricing cannot weaken the admission that relied on it. + Every plan admitted on a staged price has been priced here: staging + needs CP2, where the cheap width estimate declines and admission + prices the materialized plan (``_estimate_flat_forward``). Strictness + trades the FP32 fallback's availability for memory safety: a fused + kernel error then fails the wave, and a CP peer waits in its next + collective like after a rank-local OOM. + """ + traced = [ + self._head_backward_traced( + requests, + self._head_projection_rows( + requests, + positions=group.packed.positions_by_sequence, + lower_bound=True, + ), + ) + for group in plan.groups + if group.grad_enabled + for requests in (tuple(item.request for item in group.items),) + ] + eligible = bool(traced) and all(state is True for state in traced) + if ( + eligible + and self._plan_head_workspace_bytes(plan) + and self._checkpoint_gradient_covered( + self._plan_group_rows(plan), tuple(g.slot_ref for g in plan.groups) + ) + ): + object.__setattr__(plan, "_head_staged", True) + return eligible + + +def _plan_group_routed_rows( + self: TrainerRank, plan: _FlatForwardPlan +) -> tuple[int, ...]: + """Rows one rank's experts receive per group at balanced routing. + + HybridEP dispatches the whole EP group's rows. When that group is this + rank's CP group, a balanced rank receives the group's rows over EP, + however unevenly CP split them; otherwise keep the local rows. + """ + rows = self._plan_group_rows(plan) + if ( + not getattr(self, "_ep_group_is_cp_group", False) + or plan.signature.topology[2] <= 1 + ): + return tuple(local for local, _ in rows) + topology = self._topology() + return tuple( + min( + local, + -( + -self._cp_group_model_tokens( + _impl._pad_packed_batch(group.packed, multiple=int(topology.tp)), + topology=topology, + ) + // int(topology.cp) + ), + ) + for (local, _), group in zip(rows, plan.groups, strict=True) + ) + + def _plan_cost(self: TrainerRank, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( packed_tokens=plan.packed_tokens, @@ -543,8 +633,10 @@ def _plan_cost(self: TrainerRank, plan: _FlatForwardPlan) -> _SubforwardCost: logical_tokens=plan.active_logical_tokens, gdn_segments=plan.grad_segment_count, group_rows=self._plan_group_rows(plan), + group_routed_rows=self._plan_group_routed_rows(plan), slot_refs=tuple(g.slot_ref for g in plan.groups), head_workspace_bytes=self._plan_head_workspace_bytes(plan), + head_backward_traced=self._plan_head_backward_traced(plan), checkpoint_floor=_impl._gdn_memory.plan_floor(self, plan), retained_tokens=self._plan_retained_tokens(plan), hybridep_growth_bytes=self._plan_hybridep_growth_bytes(plan), @@ -679,11 +771,13 @@ def estimate(width: int) -> tuple[_MemoryCheck, bool, bool] | None: indices, local_inputs = local_slice(width) local_requests = list(_impl._flatten(local_inputs)) cheap_segments: list[int] = [] + cheap_traced: list[bool | None] = [] values = self._estimate_flat_forward( local_requests, checkpoint=checkpoint, sync_planning_errors=True, gdn_segments=cheap_segments, + head_traced=cheap_traced, ) if not self._all_ranks_true(values is not None): estimates[width] = None @@ -699,18 +793,27 @@ def priced( head_workspace_bytes: int, *, gdn_segments: int, + head_traced: Sequence[bool | None], + lower: bool, ) -> tuple[_MemoryCheck, int, int, _MemorySignature]: with self._planning_status(True): - required = self._estimate_required_memory_bytes_from_values( - packed_tokens=packed_tokens, - output_bytes=output_bytes, - signature=signature, - logical_tokens=logical_tokens, - # Gradient groups' segments: exact layouts' counts, else - # a bound matching the estimate's (_estimate_flat_forward). - gdn_segments=gdn_segments, - group_rows=group_rows, - head_workspace_bytes=head_workspace_bytes, + # A bound over layouts whose head may or may not stage: + # the higher price to accept, the lower to reject. + required = (min if lower else max)( + self._estimate_required_memory_bytes_from_values( + packed_tokens=packed_tokens, + output_bytes=output_bytes, + signature=signature, + logical_tokens=logical_tokens, + # Gradient groups' segments: exact layouts' counts, + # else a bound matching the estimate's + # (_estimate_flat_forward). + gdn_segments=gdn_segments, + group_rows=group_rows, + head_workspace_bytes=head_workspace_bytes, + head_backward_traced=traced, + ) + for traced in _impl._traced_states(head_traced) ) return ( self._memory_check_required(required, sync_across_dp=True), @@ -723,6 +826,7 @@ def priced_estimate( *, exact: bool, memory_minimal: bool ) -> tuple[_MemoryCheck, int, int, _MemorySignature] | None: segments: list[int] = [] + traced: list[bool | None] = [] estimated = self._estimate_flat_forward( local_requests, checkpoint=checkpoint, @@ -730,11 +834,18 @@ def priced_estimate( memory_minimal=memory_minimal, sync_planning_errors=True, gdn_segments=segments, + head_traced=traced, ) return ( None if estimated is None - else priced(*estimated, gdn_segments=sum(segments)) + else priced( + *estimated, + gdn_segments=sum(segments), + head_traced=traced, + # Only the cheap full-sharing count rejects. + lower=memory_minimal and not exact, + ) ) def trusted(packed_tokens: int, signature: _MemorySignature) -> bool: @@ -748,7 +859,12 @@ def trusted(packed_tokens: int, signature: _MemorySignature) -> bool: # reject on memory, or when it would reject on profile trust while # a profile exists — the selected layout may be far smaller than # the bound and squarely inside the profiled regime. - selected = priced(*values, gdn_segments=sum(cheap_segments)) + selected = priced( + *values, + gdn_segments=sum(cheap_segments), + head_traced=cheap_traced, + lower=False, + ) profiled = self._all_ranks_true(selected[3] in self._memory_profiles) needs_exact = not selected[0].fits or ( profiled and not trusted(selected[1], selected[3]) @@ -1415,6 +1531,8 @@ def _fill_planner_snapshot( "gdn_segments": child.grad_segment_count, "retained_tokens": self._plan_retained_tokens(child), "group_rows": self._plan_group_rows(child), + "group_routed_rows": self._plan_group_routed_rows(child), + "head_backward_traced": self._plan_head_backward_traced(child), "hybridep_growth_bytes": ( self._plan_hybridep_growth_bytes(child) ), diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index c3b7db458..68890c1c3 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -25,7 +25,12 @@ # Estimators bound on TrainerRank as plain functions, not methods. _STATIC_ESTIMATORS = frozenset( - {"_split_required_memory", "_gradient_slots", "_adapter_gradient_walk"} + { + "_split_required_memory", + "_gradient_slots", + "_adapter_gradient_walk", + "_adapter_gradient_head", + } ) @@ -112,6 +117,17 @@ def capture(rank: Any, plan: Any) -> dict[str, Any]: "_checkpoint_gradient_groups", "_checkpoint_adapter_gradient_bytes", "_adapter_gradient_walk", + "_adapter_gradient_head", + "_checkpoint_input_gradient_bytes", + "_checkpoint_gradient_covered", + "_moe_recompute_covered_for", + "_checkpoint_head_stage_bytes", + "_backward_row_state_bytes", + "_te_workspace_growth_bytes", + "_moe_checkpoint_state_bytes_per_token", + "_recomputed_mixer_bytes_per_token", + "_mixer_activation_widths", + "_triton_min_rows", ): method = getattr(rank, name) expected = getattr(_impl.TrainerRank, name) @@ -147,19 +163,19 @@ def capture(rank: Any, plan: Any) -> dict[str, Any]: head_values = _MAX_INPUT_VALUES def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: - coefficient, stages = _memory._moe_workspace_terms( + coefficient, stages, shared = _memory._moe_workspace_terms( rank, checkpoint_grad=checkpoint_grad, slot_ref=ref ) if len(stages) > 4096: raise ValueError("runtime_stage_inventory_over_limit") - reserve(64 + 72 * len(stages)) + reserve(96 + 72 * len(stages)) if ( type(coefficient) is not int or not 0 <= coefficient < 2**63 or any(v >= 2**63 for stage in stages for v in stage) ): raise ValueError("runtime_dimension_unsupported") - return [coefficient, [list(stage) for stage in stages]] + return [coefficient, [list(stage) for stage in stages], shared] has_grad = any(g.grad_enabled for g in plan.groups) for group, (physical_rows, _) in zip( @@ -260,6 +276,8 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: "head_rows": projected, "head_target_rows": target_rows, "adapter": adapter, + "moe_covered": group.grad_enabled + and rank._moe_recompute_covered_for(group.slot_ref), "gdn": None if model is None else { @@ -281,12 +299,19 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: }, } ) + layers = _memory._checkpoint_layers(rank, rank._plan_group_rows(plan)) facts = { - "version": 2, - "checkpoint_layers": _memory._checkpoint_layers( - rank, rank._plan_group_rows(plan) - ), + "version": 3, + "checkpoint_layers": layers, "checkpoint_moe_bytes_per_token": rank._checkpoint_moe_bytes_per_token(), + # Model and process readers of the recomputed layer and staged head. + "moe_checkpoint_state_bytes_per_token": ( + rank._moe_checkpoint_state_bytes_per_token() + ), + "te_workspace_growth_bytes": rank._te_workspace_growth_bytes(), + # Only a checkpointed Megatron decoder stages the head and has this. + "backward_row_state_bytes": rank._backward_row_state_bytes() if layers else 0, + "triton_min_rows": rank._triton_min_rows(), "head_vocabulary": vocabulary, "head_target_backward": target_backward, "groups": groups, @@ -314,19 +339,27 @@ def integer(value: Any, *, minimum: int = 0) -> None: "version", "checkpoint_layers", "checkpoint_moe_bytes_per_token", + "moe_checkpoint_state_bytes_per_token", + "te_workspace_growth_bytes", + "backward_row_state_bytes", + "triton_min_rows", "head_vocabulary", "head_target_backward", "groups", }, ) - if type(facts["version"]) is not int or facts["version"] != 2: + if type(facts["version"]) is not int or facts["version"] != 3: raise ValueError("unsupported runtime facts version") for key in ( "checkpoint_layers", "checkpoint_moe_bytes_per_token", + "moe_checkpoint_state_bytes_per_token", + "te_workspace_growth_bytes", + "backward_row_state_bytes", "head_vocabulary", ): integer(facts[key]) + integer(facts["triton_min_rows"], minimum=1) if type(facts["head_target_backward"]) is not bool: raise ValueError("invalid head backward eligibility") groups = facts["groups"] @@ -349,6 +382,7 @@ def integer(value: Any, *, minimum: int = 0) -> None: "head_rows", "head_target_rows", "adapter", + "moe_covered", "gdn", }, ) @@ -360,6 +394,11 @@ def integer(value: Any, *, minimum: int = 0) -> None: or len(group["slot"]) > 4096 ): raise ValueError("invalid runtime group identity") + # Only gradient groups recompute; the live reader is not asked otherwise. + if type(group["moe_covered"]) is not bool or ( + group["moe_covered"] and not group["grad"] + ): + raise ValueError("invalid MoE recompute coverage") if ( type(group["request_indices"]) is not list or len(group["request_indices"]) > 4096 @@ -381,18 +420,22 @@ def integer(value: Any, *, minimum: int = 0) -> None: terms = group[key] if ( type(terms) is not list - or len(terms) != 2 + or len(terms) != 3 or type(terms[1]) is not list or len(terms[1]) > 4096 ): raise ValueError("invalid MoE terms") - reserve(64 + 72 * len(terms[1])) + reserve(96 + 72 * len(terms[1])) integer(terms[0]) for stage in terms[1]: if type(stage) is not list or len(stage) != 2: raise ValueError("invalid MoE stage") for value in stage: integer(value) + # The shared expert's part of the coefficient (live invariant). + integer(terms[2]) + if terms[2] > terms[0]: + raise ValueError("invalid MoE terms") adapter = group["adapter"] if adapter is not None: fields(adapter, {"kind", "name", "pending"}) @@ -592,29 +635,72 @@ def verify_group( raise ValueError("head row facts disagree with selected requests/layout") def _moe_workspace_bytes( - self, rows: int, *, checkpoint_grad: bool = False, slot_ref: Any = None + self, + rows: int, + *, + routed_rows: int | None = None, + checkpoint_grad: bool = False, + slot_ref: Any = None, ) -> int: if self._facts is None: return _memory._moe_workspace_bytes( - self, rows, checkpoint_grad=checkpoint_grad, slot_ref=slot_ref + self, + rows, + routed_rows=routed_rows, + checkpoint_grad=checkpoint_grad, + slot_ref=slot_ref, ) group = self._facts["groups"][0 if slot_ref is None else slot_ref] - coefficient, stages = group["gradient" if checkpoint_grad else "forward"] + coefficient, stages, shared = group[ + "gradient" if checkpoint_grad else "forward" + ] return _memory._moe_workspace_from_terms( - rows, (coefficient, tuple(map(tuple, stages))) + rows, (coefficient, tuple(map(tuple, stages)), shared), routed_rows ) def _checkpoint_memory_floor( - self, group_rows: Any, slot_refs: Any = None, gdn_segments: int = 0 + self, + group_rows: Any, + slot_refs: Any = None, + gdn_segments: int = 0, + routed_rows: Any = None, ) -> tuple[int, int]: if self._facts is None: return _memory._checkpoint_memory_floor( - self, group_rows, slot_refs, gdn_segments + self, group_rows, slot_refs, gdn_segments, routed_rows ) return _memory._checkpoint_floor_from_facts( - self, group_rows, slot_refs, gdn_segments, self._facts["checkpoint_layers"] + self, + group_rows, + slot_refs, + gdn_segments, + self._facts["checkpoint_layers"], + routed_rows, ) + def _moe_recompute_covered_for(self, slot_ref: Any) -> bool: + if self._facts is None: + return _memory._moe_recompute_covered_for(self, slot_ref) + if slot_ref is None: + raise ValueError("replayed MoE coverage is frozen per group") + return self._facts["groups"][slot_ref]["moe_covered"] + + def _frozen(self, name: str) -> int: + assert self._facts is not None + return self._facts[name] + + def _moe_checkpoint_state_bytes_per_token(self) -> int: + return self._frozen("moe_checkpoint_state_bytes_per_token") + + def _te_workspace_growth_bytes(self) -> int: + return self._frozen("te_workspace_growth_bytes") + + def _backward_row_state_bytes(self) -> int: + return self._frozen("backward_row_state_bytes") + + def _triton_min_rows(self) -> int: + return self._frozen("triton_min_rows") + def runtime_arguments( self, facts: Any, arguments: dict[str, Any] ) -> dict[str, Any]: @@ -629,6 +715,14 @@ def runtime_arguments( raise ValueError("runtime facts disagree with selected group rows") if arguments.get("hybridep_growth_bytes", 0): raise ValueError("hybridep_runtime_facts_unsupported") + # Capture refuses expert parallelism, where each group's experts see + # its local rows (_plan_group_routed_rows). + routed = tuple(g["rows"] for g in groups) + recorded = arguments.get("group_routed_rows") + if not isinstance(recorded, (list, tuple)) or tuple(recorded) != routed: + raise ValueError("runtime facts disagree with selected routed rows") + if type(arguments.get("head_backward_traced")) is not bool: + raise ValueError("head backward staging is not recorded") head = max( max( _memory._dense_head_bytes(facts["head_vocabulary"], g["head_rows"]), @@ -660,6 +754,7 @@ def runtime_arguments( return { **arguments, "group_rows": rows, + "group_routed_rows": routed, "slot_refs": tuple(range(len(groups))), "head_workspace_bytes": head, "checkpoint_floor": (retained, workspace), diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index c0c4b1ba8..23a93a29a 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -242,6 +242,7 @@ def test_adapter_fact_validation_rejects_forged_input(change, layer, tmp_path): adapter = group["adapter"] if change == "gradient": group["grad"] = False + group["moe_covered"] = False elif change == "length": adapter["pending"] = [0] * 1026 elif change == "value": @@ -266,6 +267,95 @@ def test_adapter_fact_validation_rejects_forged_input(change, layer, tmp_path): reports.replay(report) +def test_recomputed_layer_and_head_stage_facts_are_replayed_and_frozen( + monkeypatch, tmp_path +): + rank = head_rank() + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + original, _ = emitted( + rank, rank._plan_flat_forward([request(65, grad=True)]), tmp_path + ) + item = original["replay"]["memory_replay"]["estimates"][0] + assert item["runtime_facts"]["groups"][0]["moe_covered"] is True + assert item["arguments"]["head_backward_traced"] is False + actual = reports.replay(original) + assert actual["aggregate"]["matches"] + # The replaying process's settings and TE state do not enter the answer. + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", "1") + monkeypatch.setattr(tr, "_TE_CUBLAS_WORKSPACE_BYTES", 0) + assert reports.replay(original) == actual + for field in ( + "moe_checkpoint_state_bytes_per_token", + "te_workspace_growth_bytes", + "backward_row_state_bytes", + "triton_min_rows", + ): + changed = deepcopy(original) + changed["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][field] += ( + 10**9 + ) + result = reports.replay(changed) + assert not result["estimates"][0]["matches"], field + assert ( + result["estimates"][0]["required_bytes"] + > actual["estimates"][0]["required_bytes"] + ), field + # Coverage selects the input-gradient charge and whether the head stages. + changed = deepcopy(original) + changed["replay"]["memory_replay"]["estimates"][0]["runtime_facts"]["groups"][0][ + "moe_covered" + ] = False + result = reports.replay(changed) + assert not result["estimates"][0]["matches"] + assert ( + result["estimates"][0]["required_bytes"] + != actual["estimates"][0]["required_bytes"] + ) + + +@pytest.mark.parametrize( + "change", + ["shared", "covered_no_grad", "triton_min_rows", "routed_rows", "head_staging"], +) +def test_recomputed_layer_fact_validation_rejects_forged_input(change, layer, tmp_path): + report, _, _ = adapter_report(layer, tmp_path) + item = report["replay"]["memory_replay"]["estimates"][0] + facts = item["runtime_facts"] + if change == "shared": + terms = facts["groups"][1]["gradient"] + terms[2] = terms[0] + 1 + message = "invalid MoE terms" + elif change == "covered_no_grad": + facts["groups"][0]["moe_covered"] = True + message = "invalid MoE recompute coverage" + elif change == "triton_min_rows": + facts["triton_min_rows"] = 0 + message = "invalid runtime dimension" + elif change == "routed_rows": + item["arguments"]["group_routed_rows"][1] -= 1 + message = "routed rows" + else: + del item["arguments"]["head_backward_traced"] + message = "head backward staging" + with pytest.raises(ValueError, match=message): + reports.replay(report) + + +@pytest.mark.parametrize( + "name", + ["_te_workspace_growth_bytes", "_moe_recompute_covered_for", "_triton_min_rows"], +) +def test_custom_recomputed_layer_reader_is_explicitly_incomplete(name, monkeypatch): + from art.trainer_rank import _planner_replay + + rank = _rank(monkeypatch) + plan = rank._plan_flat_forward([_request(1)]) + original = getattr(rank, name) + monkeypatch.setattr(rank, name, lambda *args: original(*args)) + with pytest.raises(ValueError, match="custom_runtime_estimator"): + _planner_replay.capture(rank, plan) + + @pytest.mark.parametrize("layers", [2**10 + 1, 0, 40.0]) def test_replay_bounds_the_recorded_layer_count(layers, pending_rank, tmp_path): rank = pending_rank diff --git a/tests/unit/test_trainer_rank_admission_inputs.py b/tests/unit/test_trainer_rank_admission_inputs.py index a79d198e1..3481bfd6d 100644 --- a/tests/unit/test_trainer_rank_admission_inputs.py +++ b/tests/unit/test_trainer_rank_admission_inputs.py @@ -32,6 +32,7 @@ def assert_plan_values(rank, plan, values): assert values["signature"] == plan.signature assert values["gdn_segments"] == plan.grad_segment_count assert values["group_rows"] == rank._plan_group_rows(plan) + assert values["group_routed_rows"] == rank._plan_group_routed_rows(plan) assert values["head_workspace_bytes"] == rank._plan_head_workspace_bytes(plan) assert values["checkpoint_floor"] == _gdn_memory.plan_floor(rank, plan) assert values["retained_tokens"] == rank._plan_retained_tokens(plan) diff --git a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py index 83c1ad50a..c27ca5e01 100644 --- a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py @@ -1,11 +1,20 @@ -"""Partial input-gradient extents: CPU admission math, not peak/overlap proof.""" +"""Input-gradient extents: CPU admission math, not peak/overlap proof. + +Where the recomputed MoE stage is priced, the floor charges the one incoming +gradient live at the last layer's recompute peak; elsewhere it keeps one +gradient per saved boundary. +""" from dataclasses import replace import pytest from test_trainer_rank_checkpoint_memory import price, rank, requests -from test_trainer_rank_moe_memory import layer # noqa: F401 -from test_trainer_rank_pending_memory import full_requests, pending_rank # noqa: F401 +from test_trainer_rank_moe_memory import _enclosing_moe, layer # noqa: F401 +from test_trainer_rank_pending_memory import ( # noqa: F401 + full_requests, + pending_rank, + rank_with_moe, +) import torch from art.trainer_rank import ForwardInput @@ -28,28 +37,93 @@ def test_pending_cold_peak_does_not_become_forward_retention(pending_rank): r = pending_rank plan = r._plan_flat_forward(full_requests()) cost = r._plan_cost(plan) - gradient = 8 * 6330 * 40 * 2048 * 2 + boundaries = 8 * 6330 * 40 * 2048 * 2 + gradient = 8 * 6330 * 2048 * 2 assert cost.checkpoint_input_gradient == gradient - # Exact previous cold estimate, including outputs and its one safety factor. + # Exact cold estimate, including outputs and its one safety factor. assert cost.retained == 23102959299 - # Unprofiled: the first execution's transients beside the workspace. assert cost.required == int( - (plan.output_bytes + 2 * gradient + 12705630112 + COLD) * 1.1 + (plan.output_bytes + boundaries + gradient + cost.checkpoint_workspace) * 1.1 ) assert r._memory_check(plan).estimated_required_bytes == cost.required profile(r, plan) warm = r._plan_cost(plan) - assert warm.retained == int((plan.output_bytes + gradient) * 1.1) - assert warm.required == int((plan.output_bytes + 2 * gradient + 12705630112) * 1.1) + assert warm.retained == int((plan.output_bytes + boundaries) * 1.1) + # Profiled: no first-execution transients. + assert warm.required == int( + (plan.output_bytes + boundaries + gradient + cost.checkpoint_workspace - COLD) + * 1.1 + ) + + +BOUNDARY_GRADIENTS = 8 * 6330 * 40 * 2048 * 2 + + +def test_single_gradient_needs_every_layer_priced(layer): + # One MoE layer among 40 dense stand-ins: the floor does not price the dense + # layers' recompute, so one gradient per boundary stays. + r = rank_with_moe(_enclosing_moe(layer), stand_in=False)[0] + assert r._moe_gradient_enclosed == (True,) and not r._moe_recompute_covered + plan = r._plan_flat_forward(full_requests()) + assert r._plan_cost(plan).checkpoint_input_gradient == BOUNDARY_GRADIENTS + + +def test_single_gradient_needs_the_fc1_stage(layer): + # Without permute fusion the FC1 stage is not enclosed: a positive FC2-only + # coefficient does not cover recompute. + moe = _enclosing_moe(layer) + moe.config.moe_permute_fusion = False + r = rank_with_moe(moe)[0] + assert r._checkpoint_moe_bytes_per_token() > 0 + assert r._moe_gradient_enclosed == (False,) and not r._moe_recompute_covered + plan = r._plan_flat_forward(full_requests()) + assert r._plan_cost(plan).checkpoint_input_gradient == BOUNDARY_GRADIENTS + + +def test_a_slot_that_loses_moe_coverage_keeps_boundary_gradients( + pending_rank, monkeypatch +): + from art.megatron.lora import LoRASlotRef + from art.trainer_rank import _impl + + r = pending_rank + ref = LoRASlotRef(kind="checkpoint", name="policy") + groups = ((100, True),) + assert r._checkpoint_input_gradient_bytes(groups) == 100 * 2048 * 2 + monkeypatch.setattr(_impl, "_moe_output_bytes_per_token", lambda *a, **k: 0) + assert r._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 40 * 4096 + + def covered(*args, enclosed, **kwargs): + enclosed.extend([True] * len(r._moe_gradient_enclosed)) + return 1 + + monkeypatch.setattr(_impl, "_moe_output_bytes_per_token", covered) + assert r._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 2048 * 2 + + +@pytest.mark.parametrize("cp", [1, 2, 4]) +def test_one_moe_gradient_only_up_to_the_traced_cp2(pending_rank, monkeypatch, cp): + """Above CP2 a rank runs more remote attention stages than the mixer's CP2 + allowance prices; the per-boundary gradient allowance must cover them.""" + r = pending_rank + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, cp, 1)) + groups = ((100, True),) + one, every = 100 * 2048 * 2, 100 * 40 * 4096 + assert r._checkpoint_input_gradient_bytes(groups) == (one if cp <= 2 else every) +@pytest.mark.parametrize("moe", [True, False]) @pytest.mark.parametrize("rows", [1, 67, 1024]) -def test_attention_only_extent_scales_with_gradient_rows(rows): +def test_attention_only_extent_scales_with_gradient_rows(rows, moe): r = rank() + if not moe: + # Without a priced MoE stage, one gradient per boundary stays: it also + # covers the dense MLP and other recompute work the floor omits. + r._moe_output_bytes_per_token = r._moe_checkpoint_grad_bytes_per_token = 0 values = r._estimate_flat_forward(requests(rows, 4096)) cost = price(r, values) assert r._gdn_layers == 0 - assert cost.checkpoint_input_gradient == rows * 40 * 2048 * 2 + assert cost.checkpoint_input_gradient == rows * (40 if not moe else 1) * 2048 * 2 assert cost.required >= int( ( cost.checkpoint_retained @@ -65,10 +139,19 @@ def test_gradient_is_not_absorbed_by_larger_head_workspace(): n, out, sig, groups, _head = r._estimate_flat_forward(requests(67, 4096)) head = 10**10 cost = price(r, (n, out, sig, groups, head)) - gradient = 67 * 40 * 2048 * 2 - assert cost.checkpoint_workspace == head + COLD - assert cost.required == int((out + head + COLD + 2 * gradient) * 1.1) - assert cost.retained == int((out + head + gradient) * 1.1) + boundaries = 67 * 40 * 2048 * 2 + gradient = 67 * 2048 * 2 + # The head stage: two more gradient-row terms, each row's RoPE and index + # state, and TE's first-GEMM workspaces. + stage = ( + head + + 2 * gradient + + 67 * r._backward_row_state_bytes() + + r._te_workspace_growth_bytes() + ) + assert cost.checkpoint_workspace == stage + COLD + assert cost.required == int((out + stage + COLD + boundaries + gradient) * 1.1) + assert cost.retained == int((out + head + boundaries) * 1.1) r._memory_profiles[sig] = _MemoryProfile( bytes_per_token=10**9, packed_tokens=n, @@ -117,11 +200,7 @@ def test_split_sums_all_gradient_children_outside_workspace_max(): profile(r, child) costs = [r._plan_cost(child) for child in children] split = _SplitForwardPlan(tuple(children), ((0,), (1,), (2,)), 3) - assert [c.checkpoint_input_gradient for c in costs] == [ - 17 * 40 * 4096, - 29 * 40 * 4096, - 0, - ] + assert [c.checkpoint_input_gradient for c in costs] == [17 * 4096, 29 * 4096, 0] expected = int( ( sum(c.checkpoint_retained + c.checkpoint_input_gradient for c in costs) @@ -163,7 +242,7 @@ def test_lower_bound_profile_cliff_preserves_separate_peak_component(): lower = r._split_chunk_lower_cost( req, tuple(q.input_tokens for q in req), checkpoint=Unset ) - assert lower.checkpoint_input_gradient == 128 * 40 * 4096 + assert lower.checkpoint_input_gradient == 128 * 4096 assert lower.required <= r._plan_cost(full).required assert lower.checkpoint_retained == full.output_bytes + 128 * 40 * 4096 @@ -228,3 +307,66 @@ def test_split_priority_subtracts_only_uncovered_gradient_peak(fully_masked): < cost.checkpoint_peak_increment < int(cost.checkpoint_input_gradient * 1.1) ) + + +def test_one_moe_gradient_still_prices_pending_adapter_gradients( + pending_rank, monkeypatch +): + from art.megatron.lora import LoRASlotRef + from art.trainer_rank import _gdn_memory + + r = pending_rank + plan = r._plan_flat_forward(full_requests()) + groups = r._plan_group_rows(plan) + retained, _ = r._checkpoint_memory_floor(groups) + # The MoE stage covers recompute: one incoming gradient, no per-boundary + # slack left to absorb gradients the backward allocates. + assert r._checkpoint_input_gradient_bytes(groups) < retained + pending = (23 * 2**20,) * 40 + (0,) + monkeypatch.setattr( + r, + "_pending_adapter_gradient_bytes", + lambda refs: pending if tuple(refs) else (), + ) + # A named slot keeping the constructor's full MoE coverage. This fixture + # has no slot tables, so price the MoE stage with the constructor's adapters. + monkeypatch.setattr(r, "_moe_recompute_covered_for", lambda ref: True) + workspace = r._moe_workspace_bytes + monkeypatch.setattr( + r, + "_moe_workspace_bytes", + lambda rows, **kwargs: workspace(rows, **{**kwargs, "slot_ref": None}), + ) + values = dict( + packed_tokens=plan.packed_tokens, + output_bytes=plan.output_bytes, + signature=plan.signature, + logical_tokens=plan.active_logical_tokens, + gdn_segments=plan.grad_segment_count, + group_rows=groups, + group_routed_rows=r._plan_group_routed_rows(plan), + head_workspace_bytes=r._plan_head_workspace_bytes(plan), + checkpoint_floor=_gdn_memory.plan_floor(r, plan), + retained_tokens=r._plan_retained_tokens(plan), + ) + slotless = (None,) * len(groups) + named = (LoRASlotRef("checkpoint", "policy"),) * len(groups) + base = r._subforward_cost(**values, slot_refs=slotless) + cost = r._subforward_cost(**values, slot_refs=named) + boundary = retained // 40 + extra = max(0, *(sum(pending[i:40]) - boundary * (39 - i) for i in range(40))) + assert base.checkpoint_adapter_gradient == 0 + assert cost.checkpoint_adapter_gradient == extra > 0 + assert cost.required == int( + ( + cost.checkpoint_retained + + cost.checkpoint_workspace + + cost.checkpoint_input_gradient + + extra + ) + * 1.1 + ) + assert ( + r._estimate_required_memory_bytes_from_values(**values, slot_refs=named) + == cost.required + ) diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index c6e33ee9d..78fcb7f73 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -8,7 +8,12 @@ import torch from art.trainer_rank import ForwardInput, TrainerRank -from art.trainer_rank._impl import Unset, _ForwardRefusal, _MemoryProfile +from art.trainer_rank._impl import ( + _TE_CUBLAS_WORKSPACE_BYTES, + Unset, + _ForwardRefusal, + _MemoryProfile, +) def rank(): @@ -53,6 +58,8 @@ def rank(): ) result._moe_output_bytes_per_token = 188416 result._moe_checkpoint_grad_bytes_per_token = 188416 + # Stands in for Qwen3.6-35B-A3B, whose MoE stage prices every layer's recompute. + result._moe_recompute_covered = True return result @@ -105,7 +112,9 @@ def test_required_and_learned_retained_use_max_not_sum(): ) cost = price(r, values) old = max(n * 2048 * 2 * 14, n * 188416) - assert cost.required == int((out + max(old, 2 * retained + work)) * 1.1) + gradient = r._checkpoint_input_gradient_bytes(groups) + assert gradient == 1024 * 2048 * 2 # The incoming gradient only. + assert cost.required == int((out + max(old, retained + gradient + work)) * 1.1) assert cost.retained == int((out + retained) * 1.1) r._memory_profiles[sig] = replace( r._memory_profiles[sig], @@ -135,7 +144,7 @@ def test_per_group_padding_precedes_gradient_filter(): assert values[0] == 40 and values[3] == ((16, True), (24, False)) assert r._checkpoint_memory_floor(values[3]) == ( 16 * 40 * 4096, - 24 * (188416 + 4 * 2048 * 2), + 24 * (188416 + 4 * 2048 * 2) + _TE_CUBLAS_WORKSPACE_BYTES, ) @@ -193,11 +202,18 @@ def test_topology_revalidated(axis): @pytest.mark.parametrize("rows", [(10, True), (11, False)]) @pytest.mark.parametrize("cp", [2, 4]) def test_cp_floor_prices_rank_rows(cp, rows): - # Callers pass rows on the most loaded CP rank; the per-row floor matches CP1. + # Callers pass rows on the most loaded CP rank; the per-row floor matches + # CP1, except that CP attention also keeps its stage buffers while a + # gradient group recomputes the layer. r = rank() single = r._checkpoint_memory_floor((rows,)) r._topology_key = lambda: (1, 1, cp, 1) - assert r._checkpoint_memory_floor((rows,)) == single != (0, 0) + retained, workspace = r._checkpoint_memory_floor((rows,)) + assert retained == single[0] and single != (0, 0) + count, grad = rows + # No attention geometry in this stub: Q and KV widths fall back to hidden. + stage = count * 2 * (3 * 2048 + 2 * 2048) if grad else 0 + assert workspace == single[1] + stage @pytest.mark.parametrize("share", [lambda n: -(-n // 2), lambda n: n * 3 // 4]) @@ -318,7 +334,8 @@ def test_split_keeps_complete_order_and_checks_each_new_subforward(): logical_per_packed=1, retained_compute_bytes_per_token=1, ) - limit = 250_000_000 + # Just below the unsplit plan's requirement: two subforwards must fit. + limit = r._plan_cost(flat).required - 1 used = 0 r._available_memory_bytes = lambda: limit - used result = r._find_admissible_forward(req, checkpoint=Unset, refusal_prefix="test") @@ -424,7 +441,8 @@ def test_no_grad_enclosure_uses_max_group_and_affine_stage(): mixed = ((3, True), (11, False)) assert r._checkpoint_memory_floor(mixed) == ( 3 * 40 * 2048 * 2, - max(3 * 188416, max(11 * 188416, 11 + 1_000_000) + 4 * 11 * 2048 * 2), + max(3 * 188416, max(11 * 188416, 11 + 1_000_000) + 4 * 11 * 2048 * 2) + + _TE_CUBLAS_WORKSPACE_BYTES, ) @@ -497,3 +515,155 @@ def test_no_grad_enclosure_config_guard(field, value): r = rank() setattr(r.runtime.model[0].decoder.config, field, value) assert r._checkpoint_memory_floor(((11, False),)) == (0, 0) + + +def test_routed_rows_move_only_the_routed_moe_part(): + # A CP2/EP2 real-data trace: the busiest CP rank held 52,480 rows, while + # HybridEP dispatched 8 x 96,794 pairs per layer across both ranks, so a + # balanced rank receives 48,397 rows' pairs. Boundaries, the mixer and the + # shared expert stay on the local rows. + r = rank() + r._topology_key = lambda: (1, 1, 2, 1) + r._moe_gradient_shared_bytes = 8192 + local, routed = 52480, 48397 + retained, workspace = r._checkpoint_memory_floor(((local, True),)) + assert r._checkpoint_memory_floor(((local, True),), None, routed_rows=(local,)) == ( + retained, + workspace, + ) + fewer = r._checkpoint_memory_floor(((local, True),), None, routed_rows=(routed,)) + assert fewer[0] == retained + assert workspace - fewer[1] == (local - routed) * (188416 - 8192) + # Never more routed rows than local ones. + assert r._moe_workspace_bytes( + 10, routed_rows=20, checkpoint_grad=True + ) == r._moe_workspace_bytes(10, checkpoint_grad=True) + r._moe_gradient_shared_bytes = 188417 + with pytest.raises(ValueError, match="shared-expert"): + r._moe_workspace_bytes(10, checkpoint_grad=True) + + +def qwen36_attention(r): + # Qwen3.6-35B-A3B attention: 16 heads and 2 query groups of 256, gated. + r._geometry = replace( + r._geometry, num_attention_heads=16, num_query_groups=2, kv_channels=256 + ) + r._attention_output_gate = True + return r + + +@pytest.mark.parametrize( + "cp,expected", + # Traces measured 66 KB per token at CP1 and 94 KB at CP2. CP4 reuses the + # CP2 stage allowance; it is not measured. + [(1, 2 * (2 * 2048 + 7 * 4096 + 3 * 512)), (2, 95232), (4, 95232)], +) +def test_recomputed_attention_is_priced_beside_the_moe_stage(cp, expected): + r = qwen36_attention(rank()) + r._topology_key = lambda: (1, 1, cp, 1) + assert r._recomputed_mixer_bytes_per_token() == expected + moe = r._moe_workspace_bytes(10, checkpoint_grad=True) + # Beside the mixer: the recomputed layer's residual and pre-MLP norm rows, + # and Transformer Engine's cuBLAS workspaces. + assert ( + r._checkpoint_memory_floor(((10, True),))[1] + == moe + 10 * (expected + 2 * 2048 * 2) + _TE_CUBLAS_WORKSPACE_BYTES + ) + + +@pytest.mark.parametrize( + "cp,gdn_layers,expected", + [ + # Hybrid: GDN (82 KB) is larger at CP1, CP attention (95 KB) at CP2. + (1, 30, 81920), + (2, 30, 95232), + # GDN only: CP exchanges add a hidden-width and a value-width row. + (1, 40, 81920), + (2, 40, 94208), + ], +) +def test_larger_recomputed_mixer_is_priced(cp, gdn_layers, expected): + r = qwen36_attention(rank()) + r._geometry = replace( + r._geometry, + gdn_key_heads=16, + gdn_key_head_dim=128, + gdn_value_heads=32, + gdn_value_head_dim=128, + ) + r._gdn_layers = gdn_layers + r._topology_key = lambda: (1, 1, cp, 1) + assert r._recomputed_mixer_bytes_per_token() == expected + + +def test_no_grad_groups_do_not_recompute_a_mixer(): + r = qwen36_attention(rank()) + moe = r._moe_workspace_bytes(10) + assert r._checkpoint_memory_floor(((10, False),)) == (0, moe + 4 * 10 * 2048 * 2) + + +def test_ungated_attention_prices_fewer_projections(): + r = qwen36_attention(rank()) + r._attention_output_gate = False + assert r._recomputed_mixer_bytes_per_token() == 2 * (2 * 2048 + 5 * 4096 + 3 * 512) + r._topology_key = lambda: (1, 1, 2, 1) + assert r._recomputed_mixer_bytes_per_token() == 2 * ( + 2 * 2048 + 5 * 4096 + 3 * 512 + 3 * 4096 + 2 * 512 + ) + + +@pytest.mark.parametrize("cp", [1, 2]) +def test_gdn_width_follows_hidden_key_and_value_separately(cp): + # Hidden differs from the key width, as in Qwen3.5-27B: CP exchanges + # carry hidden-width inputs and value-width outputs. + r = rank() + r._hidden_size = r.runtime.model[0].decoder.config.hidden_size = 5120 + r._geometry = replace( + r._geometry, + gdn_key_heads=16, + gdn_key_head_dim=128, + gdn_value_heads=48, + gdn_value_head_dim=128, + ) + r._gdn_layers = r._num_layers + r._topology_key = lambda: (1, 1, cp, 1) + key, value = 16 * 128, 48 * 128 + # q and k are l2-normalized after expansion to the 48 value heads. + width = 5120 + 2 * key + 2 * 48 * 128 + 6 * value + 64 * 48 + if cp > 1: + width += 5120 + value + assert r._recomputed_mixer_bytes_per_token() == 2 * width + moe = r._moe_workspace_bytes(7, checkpoint_grad=True) + assert ( + r._checkpoint_memory_floor(((7, True),))[1] + == moe + 7 * 2 * (width + 2 * 5120) + _TE_CUBLAS_WORKSPACE_BYTES + ) + + +@pytest.mark.parametrize("ep,routing", [(1, 3332), (2, 3600)]) +def test_moe_state_beside_the_recomputed_stage(ep, routing): + # Qwen3.6-35B-A3B: 256 experts, 512-wide shared expert. Traces measured + # 3,332 (EP1) and 3,593 (EP2) bytes of routing state per local token. + r = rank() + r._geometry = replace(r._geometry, moe_experts=256, moe_shared_expert_ffn=512) + r._parallel_shape = replace(r._parallel_shape, ep=ep) + assert r._moe_checkpoint_state_bytes_per_token() == routing + 3 * 512 * 2 + r._geometry = replace(r._geometry, moe_experts=0) + assert r._moe_checkpoint_state_bytes_per_token() == 0 + + +def test_te_workspaces_are_growth_until_allocated(monkeypatch): + gemm = pytest.importorskip("transformer_engine.pytorch.cpp_extensions.gemm") + r = rank() + entries = [0] + + class Cached: + def cache_info(self): + return SimpleNamespace(currsize=entries[0]) + + monkeypatch.setattr(gemm, "get_cublas_workspace", Cached()) + assert r._te_workspace_growth_bytes() == _TE_CUBLAS_WORKSPACE_BYTES + entries[0] = 1 # Plain GEMM only: the grouped streams are still to come. + assert r._te_workspace_growth_bytes() == _TE_CUBLAS_WORKSPACE_BYTES + entries[0] = 2 + assert r._te_workspace_growth_bytes() == 0 diff --git a/tests/unit/test_trainer_rank_converted_memory.py b/tests/unit/test_trainer_rank_converted_memory.py index a9c809163..fe1b47372 100644 --- a/tests/unit/test_trainer_rank_converted_memory.py +++ b/tests/unit/test_trainer_rank_converted_memory.py @@ -1,16 +1,21 @@ """Source-derived affine routed-expert stages; no complete backward/compiled bound.""" +import math from types import SimpleNamespace from typing import Any, cast import pytest -from test_trainer_rank_moe_memory import _enclosing_moe, _rank +from test_trainer_rank_moe_memory import _enclosing_moe, _hybridep, _rank from test_trainer_rank_moe_memory import layer as layer from test_trainer_rank_pending_memory import module, rank_with_moe import torch from art.trainer_rank import ForwardInput, _gdn_memory -from art.trainer_rank._impl import _expert_lora_weight_storage +from art.trainer_rank._impl import ( + _expert_lora_weight_storage, + _moe_output_bytes_per_token, +) +from art.trainer_rank._planner_cost import ParallelShape def weights(layer: Any, rank: int, *, fc1: bool = True, dtype=torch.bfloat16): @@ -96,9 +101,19 @@ def test_actual_plan_cost_and_admission(layer, rank_value, grad, output): retained, workspace = rank._checkpoint_memory_floor(rank._plan_group_rows(plan)) pending = _gdn_memory.plan_floor(rank, plan) if grad: - assert workspace == expected(8, rank_value, True) + # The recomputed layer's mixer, its residual and pre-MLP norm rows and + # its MoE routing state stay live beside its MoE stage; the first call + # also allocates TE's cuBLAS workspaces. + mixer = rank._recomputed_mixer_bytes_per_token() + beside = 2 * 2048 * 2 + rank._moe_checkpoint_state_bytes_per_token() + assert mixer > 0 + assert workspace == ( + expected(8, rank_value, True) + + 8 * (mixer + beside) + + rank._te_workspace_growth_bytes() + ) + # The GDN pending floor combines with this one by maximum. assert pending[0] == retained == 8 * 40 * 2048 * 2 - assert pending[1] >= workspace else: assert retained == 0 and pending == (0, 0) assert workspace == expected(8, rank_value, False) + 4 * 8 * 2048 * 2 @@ -124,8 +139,15 @@ def test_reference_and_gradient_keep_distinct_stage_modes(layer, order, rank_val assert set(groups) == {(3, True), (9, False)} retained, workspace = rank._checkpoint_memory_floor(groups) assert retained == 3 * 40 * 2048 * 2 - assert workspace == max( - expected(3, rank_value, True), expected(9, rank_value, False) + 4 * 9 * 2048 * 2 + beside = 2 * 2048 * 2 + rank._moe_checkpoint_state_bytes_per_token() + assert ( + workspace + == max( + expected(3, rank_value, True) + + 3 * (rank._recomputed_mixer_bytes_per_token() + beside), + expected(9, rank_value, False) + 4 * 9 * 2048 * 2, + ) + + rank._te_workspace_growth_bytes() ) assert ( rank._memory_check(plan).estimated_required_bytes @@ -262,6 +284,38 @@ def test_fc1_fixed_weights_do_not_scale_with_topk(layer, topk): ) +def test_hybridep_fc1_stages_hold_one_dispatched_input(layer): + # FC1's converted stages hold the routed H-wide inputs too: two under the + # EP1 all-to-all, one under HybridEP, over its EP2 allowance rows. + weights(layer, 8) + single: list[tuple[int, int]] = [] + _moe_output_bytes_per_token( + [layer], + ParallelShape(tp=1, cp=1), + checkpoint_grad=True, + converted_stages=single, + ) + expert = _hybridep(layer, 2) + expert.token_dispatcher.num_local_experts = 128 + sharded: list[tuple[int, int]] = [] + _moe_output_bytes_per_token( + [expert], + ParallelShape(tp=1, cp=2, ep=2), + checkpoint_grad=True, + converted_stages=sharded, + ) + assert [stage[0] for stage in single[:2]] == [ + 16 * (2 * 2048 + 2 * 1024 + 8), + 16 * (2 * 2048 + 3 * 1024 + 8), + ] + # 8 x 1.4 routed rows at EP2, rounded up per stage. + assert [stage[0] for stage in sharded[:2]] == [ + math.ceil(8 * 1.4 * 2 * (2048 + 2 * 1024 + 8)), + math.ceil(8 * 1.4 * 2 * (2048 + 3 * 1024 + 8)), + ] + assert [stage[1] for stage in sharded] == [stage[1] for stage in single] + + def test_wide_fc1_sum_is_a_separate_stage(layer): weights(layer, 8) experts = layer.experts diff --git a/tests/unit/test_trainer_rank_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py index 54d279f91..569037e40 100644 --- a/tests/unit/test_trainer_rank_head_memory.py +++ b/tests/unit/test_trainer_rank_head_memory.py @@ -146,13 +146,18 @@ def test_outputs_retention_and_empirical_peak_are_counted_once(): r = rank() plan = r._plan_flat_forward([request(grad=True)]) retained = 512 * 40 * 2048 * 2 - gradient = 512 * 40 * 2048 * 2 + gradient = 512 * 2048 * 2 # The incoming gradient beside the MoE stage. head = 3 * 512 * 248320 * 2 cost = r._plan_cost(plan) assert cost.retained == int((plan.output_bytes + retained + head) * 1.1) - # Unprofiled: the first execution's transients beside the head workspace. + # Unprofiled head stage: the head workspace, the final output, saved + # selected rows and hidden-row gradient beside the incoming one, each + # row's RoPE and index state, TE's first-GEMM workspaces and the first + # execution's transients. + state = 512 * r._backward_row_state_bytes() + te = r._te_workspace_growth_bytes() assert cost.required == int( - (plan.output_bytes + retained + gradient + head + COLD) * 1.1 + (plan.output_bytes + retained + 3 * gradient + state + head + te + COLD) * 1.1 ) r._memory_profiles[plan.signature] = _MemoryProfile( bytes_per_token=2_000_000, @@ -312,14 +317,21 @@ def test_tied_standard_head_weight_uses_the_same_capacity(): @pytest.mark.parametrize("rows", [128, 512]) -def test_target_backward_refuses_budget_below_logits_and_both_gradients(rows): +def test_target_backward_refuses_budget_below_logits_and_both_gradients( + monkeypatch, rows +): r = rank() + # Isolate the head term from TE's one-time cuBLAS workspace growth, and + # the three target-backward buffers from the small-chunk fallback cover. + r._te_workspace_growth_bytes = lambda: 0 + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", "1") plan = r._plan_flat_forward([request(rows, grad=True)]) retained, _ = r._checkpoint_memory_floor(r._plan_group_rows(plan)) - gradient = rows * 40 * 2048 * 2 + gradient = rows * 2048 * 2 + stage = 3 * gradient + rows * r._backward_row_state_bytes() dense = min(rows, 512) * 248320 * 2 - before = int((plan.output_bytes + retained + gradient + 2 * dense + COLD) * 1.1) - expected = int((plan.output_bytes + retained + gradient + 3 * dense + COLD) * 1.1) + before = int((plan.output_bytes + retained + stage + 2 * dense + COLD) * 1.1) + expected = int((plan.output_bytes + retained + stage + 3 * dense + COLD) * 1.1) r._available_memory_bytes = lambda: (before + expected) // 2 check = r._memory_check(plan) assert check.estimated_required_bytes == expected diff --git a/tests/unit/test_trainer_rank_head_stage_memory.py b/tests/unit/test_trainer_rank_head_stage_memory.py new file mode 100644 index 000000000..695ebe274 --- /dev/null +++ b/tests/unit/test_trainer_rank_head_stage_memory.py @@ -0,0 +1,456 @@ +"""The head's backward and the decoder's recompute peak apart: CPU admission math. + +Qwen3.6-35B-A3B CP2 allocator traces (EP1 and EP2): the head's buffers are +freed before the recompute backward allocates its workspace and adapter +gradients, and TE's first-GEMM workspaces are live at the head's peak. +""" + +from dataclasses import replace +import itertools +import random +from typing import Any + +import pytest +from test_trainer_rank_adapter_gradient_memory import ( + OTHER, + POLICY, + oracle, + priced, + with_pending, + with_slot_pending, +) +from test_trainer_rank_checkpoint_memory import rank, requests + +from art.megatron.lora import LoRASlotRef +from art.trainer_rank import TrainerRank +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) +from art.trainer_rank._impl import _MemoryProfile + +MiB = 2**20 + + +def covered_rank(monkeypatch): + r = rank() + # The MoE stage covers the policy slot's recompute, as the constructor's. + monkeypatch.setattr(r, "_moe_recompute_covered_for", lambda ref: True) + return r + + +def head_oracle(chains): + """Worst adapter bytes live beyond the floor while one group's head runs. + + Any set of the other groups may have run their backward first: each holds + all its gradients and none of its boundaries. The running group holds only + its gradients outside the decoder. + """ + worst = 0 + for index, (pending, _) in enumerate(chains): + others = chains[:index] + chains[index + 1 :] + for count in range(len(others) + 1): + for done in itertools.combinations(others, count): + live = pending[-1] + sum(sum(p) - sum(b) for p, b in done) + worst = max(worst, live) + return worst + + +def test_head_adapter_gradients_match_every_order_for_random_sizes(monkeypatch): + r = rank() + generator = random.Random(2) + slots = [LoRASlotRef("checkpoint", f"slot{index}") for index in range(4)] + for _ in range(200): + layers = generator.randint(1, 5) + chains, groups, by_slot = [], [], {} + for slot in slots[: generator.randint(1, 4)]: + pending = [generator.randint(0, 30) for _ in range(layers + 1)] + boundaries = [generator.randint(0, 30) for _ in range(layers)] + by_slot[(slot,)] = pending + chains.append((pending, boundaries)) + groups.append((slot, boundaries)) + with_slot_pending(monkeypatch, r, by_slot) + assert r._checkpoint_adapter_gradient_bytes(groups, head=True) == head_oracle( + chains + ) + + +def test_one_group_head_meets_only_its_gradients_outside_the_decoder(monkeypatch): + r = rank() + with_pending(monkeypatch, r, [30] * 4 + [7]) + assert r._checkpoint_adapter_gradient_bytes(((POLICY, [2] * 4),), head=True) == 7 + # The decoder stage still meets its own gradients. + assert r._checkpoint_adapter_gradient_bytes(((POLICY, [2] * 4),)) == oracle( + [30] * 4 + [7], [2] * 4 + ) + with_slot_pending( + monkeypatch, r, {(POLICY,): [30] * 4 + [7], (OTHER,): [1] * 4 + [0]} + ) + # The policy group may run first and raise the other's head; a group whose + # boundaries outweigh its gradients never raises the policy's. + groups = ((POLICY, [2] * 4), (OTHER, [40] * 4)) + assert r._checkpoint_adapter_gradient_bytes(groups, head=True) == 127 - 8 + groups = ((POLICY, [2] * 4), (OTHER, [2] * 4)) + with_slot_pending( + monkeypatch, r, {(POLICY,): [30] * 4 + [7], (OTHER,): [0] * 4 + [0]} + ) + assert r._checkpoint_adapter_gradient_bytes(groups, head=True) == 127 - 8 + + +def staged(r, values, slot_refs): + """``priced`` for a head whose backward is the traced one.""" + n, out, signature, groups, head = values + return r._subforward_cost( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=slot_refs, + head_workspace_bytes=head, + head_backward_traced=True, + ) + + +@pytest.mark.parametrize("head", [0, 64 * MiB, 700 * MiB, 3000 * MiB]) +def test_cost_and_estimate_price_the_larger_stage(monkeypatch, head): + r = covered_rank(monkeypatch) + n, out, signature, groups, _ = r._estimate_flat_forward(requests(67, 4096)) + values = (n, out, signature, groups, head) + retained, workspace = r._checkpoint_memory_floor(groups) + gradient = 67 * 2048 * 2 + boundary = retained // 40 + pending = [23 * MiB] * 40 + [0] + with_pending(monkeypatch, r, pending) + extra = oracle(pending, [boundary] * 40) + te = r._te_workspace_growth_bytes() + state = 67 * r._backward_row_state_bytes() + decoder = workspace + extra + stage = head + 2 * gradient + state + te if head else 0 + cost = staged(r, values, (POLICY, None)) + assert cost.checkpoint_input_gradient == gradient + assert cost.required == int( + (out + retained + gradient + max(decoder, stage) + COLD) * 1.1 + ) + estimate = r._estimate_required_memory_bytes_from_values( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=(POLICY, None), + head_workspace_bytes=head, + head_backward_traced=True, + ) + assert estimate == cost.required + # An untraced head keeps the unstaged price as a floor. + unstaged = max(workspace, head) + extra + untraced = priced(r, values, (POLICY, None)) + assert untraced.required == int( + (out + retained + gradient + max(unstaged, decoder, stage) + COLD) * 1.1 + ) + assert untraced.checkpoint_workspace == cost.checkpoint_workspace + # Warm: no first-execution transients and no TE growth in either stage. + r._te_workspace_growth_bytes = lambda: 0 + r._memory_profiles[signature] = _MemoryProfile(bytes_per_token=1, packed_tokens=n) + decoder = r._checkpoint_memory_floor(groups)[1] + extra + assert decoder == workspace - te + extra + stage = head + 2 * gradient + state if head else 0 + assert staged(r, values, (POLICY, None)).required == int( + (out + retained + gradient + max(decoder, stage)) * 1.1 + ) + + +def test_head_bound_short_wave_no_longer_adds_the_adapter_extra(monkeypatch): + r = covered_rank(monkeypatch) + n, out, signature, groups, _ = r._estimate_flat_forward(requests(67, 4096)) + retained, workspace = r._checkpoint_memory_floor(groups) + head = 2 * workspace + pending = [23 * MiB] * 40 + [0] + with_pending(monkeypatch, r, pending) + extra = oracle(pending, [retained // 40] * 40) + assert head > workspace and extra > 800 * MiB + values = (n, out, signature, groups, head) + unstaged = int((out + retained + head + 67 * 2048 * 2 + extra + COLD) * 1.1) + assert staged(r, values, (POLICY, None)).required < unstaged + # Untraced (e.g. a top-k-only head, whose backward keeps its recomputed + # logits beside their gradient): never below the unstaged price. + assert priced(r, values, (POLICY, None)).required == unstaged + + +@pytest.mark.parametrize("uncovered", ["cp4", "uncovered_moe"]) +def test_untraced_recompute_keeps_the_head_in_the_decoder_stage(monkeypatch, uncovered): + r = covered_rank(monkeypatch) + n, out, signature, groups, _ = r._estimate_flat_forward(requests(67, 4096)) + if uncovered == "cp4": + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 4, 1)) + else: + r._moe_output_bytes_per_token = r._moe_checkpoint_grad_bytes_per_token = 0 + head = 700 * MiB + retained, workspace = r._checkpoint_memory_floor(groups) + gradient = r._checkpoint_input_gradient_bytes(groups) + assert gradient == retained # One gradient per boundary. + with_pending(monkeypatch, r, ()) + # Even for a traced head. + cost = staged(r, (n, out, signature, groups, head), (POLICY, None)) + assert cost.checkpoint_workspace == max(workspace, head) + COLD + assert cost.required == int( + (out + retained + max(workspace, head) + gradient + COLD) * 1.1 + ) + + +def test_split_keeps_each_head_stage_beside_every_adapter_gradient(monkeypatch): + r = covered_rank(monkeypatch) + head = 700 * MiB + values = r._estimate_flat_forward(requests(67, 4096)) + n, out, signature, groups, _ = values + pending = [23 * MiB] * 40 + [0] + with_pending(monkeypatch, r, pending) + left = staged(r, (n, out, signature, groups, head), (POLICY, None)) + right = staged(r, (n, out, signature, groups, head), (OTHER, None)) + gradient = left.checkpoint_input_gradient + stage = ( + head + + 2 * gradient + + 67 * r._backward_row_state_bytes() + + r._te_workspace_growth_bytes() + ) + assert ( + left.checkpoint_workspace + == max(r._checkpoint_memory_floor(groups)[1], stage) + COLD + ) + # The left child's head can run after the right child's decoder backward: + # both children's boundaries and incoming gradients, the left head stage + # and every adapter gradient the right one allocated (distinct slots). + after_right = ( + left.checkpoint_retained + + right.checkpoint_retained + + 2 * gradient + + stage + + COLD + + right.checkpoint_adapter_gradient + ) + split = TrainerRank._split_required_memory([left, right]) + assert split >= int(after_right * 1.1) + assert split >= int( + (after_right + left.checkpoint_adapter_gradient) * 1.1 + ) # Either may run first, and distinct slots each allocate their own. + + +def test_traced_head_backward_is_target_only_on_the_traced_path(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + from art.trainer_rank import _impl + + r = head_rank() + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + target = request(512, grad=True) + # Until the fused statistics have run in this process, nothing is traced. + state = {"succeeded": set(), "failed": False} + monkeypatch.setattr(_impl, "_TRITON_STATS_STATE", state) + assert r._head_backward_traced([target], 512) is False + # The top-k kernel's success does not prove the target-only one. + state["succeeded"].add("local_topk_stats") + assert r._head_backward_traced([target], 512) is False + state["succeeded"].add("local_logsumexp_stats") + assert r._head_backward_traced([target], 512) is True + # Top-k, logits and hidden-state outputs keep further dense gradients. + for extra in ({"top_k": 2}, {"logits": True}, {"hidden_states": True}): + other = replace(target, **extra) + assert r._head_backward_traced([target, other], 512) is False + topk_only = replace(target, target_tokens=None, top_k=2) + assert r._head_backward_traced([topk_only], 512) is False + # The fused statistics need 64 rows; bounds that straddle them are open. + assert r._head_backward_traced([target], 63) is False + assert r._head_backward_traced([target], 63, 64) is None + assert r._head_backward_traced([target], 32, 63) is False + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_TOPK", "0") + assert r._head_backward_traced([target], 512) is False + monkeypatch.delenv("ART_TRAINER_RANK_TRITON_TOPK") + # Only CP2 was traced. + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 1, 1)) + assert r._head_backward_traced([target], 512) is False + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + # One error sent a kernel to the FP32 fallback silently: never again. + state["failed"] = True + assert r._head_backward_traced([target], 512) is False + state["failed"] = False + r.runtime.model[0].config.use_mup = True + assert r._head_backward_traced([target], 512) is False + + +def test_undecided_heads_bound_both_ways(): + from art.trainer_rank._impl import _traced_states + + assert _traced_states([]) == (False,) + assert _traced_states([True, True]) == (True,) + assert _traced_states([True, False, None]) == (False,) + # Lower bounds take the cheaper, acceptance the dearer of both. + assert _traced_states([True, None]) == (True, False) + + +def test_each_row_keeps_its_rope_embedding_and_index_state(): + import torch + + r = rank() + model = r.runtime.model[0] + assert r._backward_row_state_bytes() == 256 + # Qwen3.6-35B-A3B: a 64-wide rotary embedding (32 frequencies), FP32. + model.rotary_pos_emb = torch.nn.Module() + model.rotary_pos_emb.inv_freq = torch.ones(32) + assert r._backward_row_state_bytes() == 64 * 4 + 256 + + +def test_plan_stages_only_traced_gradient_heads(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + from art.trainer_rank import _impl + + monkeypatch.setattr( + _impl, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, + ) + r = head_rank() + target = request(512, grad=True) + topk_only = replace(target, target_tokens=None, top_k=2) + plans = [r._plan_flat_forward([request]) for request in (target, topk_only)] + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + # A top-k-only head's backward keeps its recomputed logits beside their + # gradient, twice the one buffer its workspace prices: never staged. + assert [r._plan_head_backward_traced(plan) for plan in plans] == [True, False] + no_grad = r._plan_flat_forward([request(512)]) + assert r._plan_head_backward_traced(no_grad) is False + + +def test_head_stage_covers_a_small_chunks_fp32_fallback(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + + r = head_rank() + monkeypatch.setattr(r, "_moe_recompute_covered_for", lambda ref: True) + monkeypatch.setattr(r, "_checkpoint_gradient_covered", lambda *a, **k: True) + r._te_workspace_growth_bytes = lambda: 0 + groups = ((8, True),) + dense = 248320 * 2 + state = 8 * r._backward_row_state_bytes() + # However the CP split leaves a rank's projected rows, a chunk below the + # fused minimum takes the FP32 fallback: nine buffers of up to 63 rows. + stage = r._checkpoint_head_stage_bytes(3 * 8 * dense, 0, groups, None) + assert stage == 9 * 63 * dense + state + stage = r._checkpoint_head_stage_bytes(3 * 512 * dense, 0, groups, None) + assert stage == 3 * 512 * dense + state + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", "512") + stage = r._checkpoint_head_stage_bytes(3 * 512 * dense, 0, groups, None) + assert stage == 9 * 511 * dense + state + + +def test_a_staged_plan_runs_its_fused_statistics_strictly(monkeypatch): + from types import SimpleNamespace + + from art.trainer_rank import _impl, topk + + state = {"succeeded": {"local_logsumexp_stats"}, "failed": False} + monkeypatch.setattr(_impl, "_TRITON_STATS_STATE", state) + + def fail(*args, **kwargs): + raise RuntimeError("kernel launch failed") + + monkeypatch.setattr(topk, "local_logsumexp_stats", fail) + chunk: Any = SimpleNamespace(is_cuda=True, shape=(512, 248320)) + # Unstaged: the FP32 fallback, and no later plan stages again. + assert _impl._try_triton_stats("local_logsumexp_stats", chunk) is None + assert state["failed"] is True + # Staged: the plan's price assumed the fused path, so it raises instead. + with pytest.raises(RuntimeError, match="admitted on their memory"): + _impl._try_triton_stats("local_logsumexp_stats", chunk, strict=True) + # Too few rows is a predictable fallback the head stage prices. + small = SimpleNamespace(is_cuda=True, shape=(63, 248320)) + assert _impl._try_triton_stats("local_logsumexp_stats", small, strict=True) is None + + +def test_execution_binds_strictness_to_the_staged_admission(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + from art.trainer_rank import ForwardOutput, _impl + + monkeypatch.setattr( + _impl, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, + ) + r = head_rank() + monkeypatch.setattr(r, "_moe_recompute_covered_for", lambda ref: True) + plan = r._plan_flat_forward([request(512, grad=True), request(16, hidden=True)]) + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + monkeypatch.setattr(r, "_topology", lambda: None) + monkeypatch.setattr(r, "_validate_hybridep_topology", lambda: None) + monkeypatch.setattr(r, "_configure_hybridep", lambda *a, **k: None) + monkeypatch.setattr(r, "_prepare_packed_forward", lambda packed: None) + seen = [] + + def forward_packed(items, prepared): + seen.append(_impl._HEAD_STATISTICS_STRICT.get()) + return [ForwardOutput(None, None, None, None)] * len(items) + + monkeypatch.setattr(r, "_forward_packed", forward_packed) + r._execute_flat_plan(plan) + assert seen == [False, False] # Never priced staged. + assert r._plan_head_backward_traced(plan) is True + assert plan._head_staged is True + seen.clear() + r._execute_flat_plan(plan) + # Only the gradient group runs strictly; nothing leaks past execution. + grad_first = [group.grad_enabled for group in plan.groups] + assert seen == grad_first + assert _impl._HEAD_STATISTICS_STRICT.get() is False + # A later price that no longer stages (a kernel failed since) cannot + # weaken the admission that relied on it. + _impl._TRITON_STATS_STATE["failed"] = True + assert r._plan_head_backward_traced(plan) is False + assert plan._head_staged is True + + +def test_an_eligible_head_the_price_does_not_stage_is_not_strict(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + from art.trainer_rank import _impl + + monkeypatch.setattr( + _impl, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, + ) + r = head_rank() + plan = r._plan_flat_forward([request(512, grad=True)]) + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + # The recompute is not covered: the head keeps the unstaged price. + monkeypatch.setattr(r, "_checkpoint_gradient_covered", lambda *a, **k: False) + assert r._plan_head_backward_traced(plan) is True + assert getattr(plan, "_head_staged", False) is False + + +def test_several_labels_per_row_are_not_traced(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + from art.trainer_rank import _impl + + monkeypatch.setattr( + _impl, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, + ) + r = head_rank() + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + single = request(512, grad=True) + several = replace( + single, target_tokens=single.target_tokens.unsqueeze(1).repeat(1, 4) + ) + assert r._head_backward_traced([single], 512) is True + assert r._head_backward_traced([several], 512) is False + # One label per token over a leading batch axis is still one per row. + batched = replace(single, input_tokens=single.input_tokens.unsqueeze(0)) + assert r._head_backward_traced([batched], 512) is True diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 9cb0315ba..7bc27d8fd 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -1,6 +1,7 @@ """CPU contracts for one known MoE component, not a whole-model memory bound.""" from dataclasses import replace +import math from types import SimpleNamespace from typing import Any, cast import weakref @@ -14,6 +15,7 @@ ) from art.trainer_rank._impl import ( _PACKED_PRICED_LOGICAL_ROW_BYTES, + _ep_routed_row_allowance, _MemoryProfile, _MemorySignature, _moe_output_bytes_per_token, @@ -206,20 +208,33 @@ def _hybridep(layer, ep: int, manager: str = "hybridep"): return layer -@pytest.mark.parametrize("ep", [2, 4]) -def test_hybridep_prices_routed_rows_with_imbalance_allowance(layer, ep): +@pytest.mark.parametrize("ep,allowance", [(2, 1.4), (4, 1.6), (8, 2.0), (16, 2.4)]) +def test_hybridep_prices_routed_rows_with_imbalance_allowance(layer, ep, allowance): # HybridEP gives each rank the pairs routed to its local experts: balanced - # routing matches EP1's top-k rows per local token, with a 1.5x allowance. + # routing matches EP1's top-k rows per local token, scaled by the measured + # EP-dependent worst-layer imbalance (log2 growth beyond EP8). single = _moe_output_bytes_per_token([layer], ParallelShape(tp=1, cp=1)) sharded = ParallelShape(tp=1, cp=ep, ep=ep) expert = _moe_output_bytes_per_token([_hybridep(layer, ep)], sharded) - # Top-k 8 routed rows per local token become 12; no shared experts here. - assert single > 0 and expert == single // 8 * 12 + # Top-k 8 routed rows per local token become 8 x allowance; no shared + # experts here. + assert _ep_routed_row_allowance(ep) == pytest.approx(allowance) + assert single > 0 and expert == math.ceil(single * allowance) + + +def test_unmeasured_ep_sizes_use_the_next_measured_allowance(): + assert _ep_routed_row_allowance(1) == 1.0 + assert _ep_routed_row_allowance(3) == _ep_routed_row_allowance(4) == 1.6 + assert _ep_routed_row_allowance(6) == 2.0 + # Never above every pair on one rank. + assert _ep_routed_row_allowance(1024) <= 1024 def test_hybridep_keeps_the_enclosing_fc1_stage(layer): - # The FC1 inputs and gate/up sum stay live at the FC2 sum under HybridEP - # too; EP1's two dispatched H-wide inputs over-count HybridEP's one. + # The FC1 input and gate/up sum stay live at the FC2 sum under HybridEP + # too. The EP1 all-to-all path holds two routed H-wide inputs (its + # permuted rows and their expert-sorted copy); HybridEP permutes while + # dispatching and holds one, as Qwen3.6 CP2 allocator traces show. single = _moe_output_bytes_per_token( [_enclosing_moe(layer)], ParallelShape(tp=1, cp=1) ) @@ -227,7 +242,7 @@ def test_hybridep_keeps_the_enclosing_fc1_stage(layer): expert.token_dispatcher.num_local_experts = 128 sharded = _moe_output_bytes_per_token([expert], ParallelShape(tp=1, cp=2, ep=2)) assert single == 8 * (512 + 3 * 2048 + 2 * 2048 + 1024) * 2 - assert sharded == single // 8 * 12 + assert sharded == math.ceil(8 * 1.4 * (512 + 3 * 2048 + 2048 + 1024) * 2) @pytest.mark.parametrize( @@ -594,7 +609,7 @@ def held(rows): monkeypatch.setattr( rank, "_checkpoint_memory_floor", - lambda rows, refs=None, segments=0: (10**7, 10**6), + lambda rows, refs=None, segments=0, routed_rows=None: (10**7, 10**6), ) grown = rank._plan_cost(plan) monkeypatch.setattr(rank, "_plan_hybridep_growth_bytes", lambda plan: 0) @@ -613,6 +628,92 @@ def held(rows): assert rank._plan_hybridep_growth_bytes(plan) == 0 +def test_converted_stages_follow_routed_rows_and_shared_stays_local(): + rank = _rank() + rank._moe_output_bytes_per_token = 1000 + rank._moe_forward_stages = ((1200, 50),) + rank._moe_forward_shared_bytes = 100 + assert rank._moe_workspace_bytes(10) == 10 * 1200 + 50 + assert rank._moe_workspace_bytes(10, routed_rows=6) == 6 * 1200 + 50 + 4 * 100 + assert rank._moe_workspace_bytes(10, routed_rows=0) == 50 + 10 * 100 + + +def test_ep_group_must_be_exactly_the_cp_group(monkeypatch): + ps = pytest.importorskip("megatron.core.parallel_state") + from art.trainer_rank import _impl + + shape = ParallelShape(tp=1, cp=2, ep=2) + for other in ( + replace(shape, ep=1), + replace(shape, ep=4), + replace(shape, tp=2), + replace(shape, etp=2), + ): + assert not _impl._ep_group_is_cp_group(other) + # Without initialized process groups there is nothing to compare. + assert not _impl._ep_group_is_cp_group(shape) + monkeypatch.setattr(_impl.dist, "is_initialized", lambda: True) + expert, context = object(), object() + ranks = {id(expert): [0, 1], id(context): [1, 0]} + monkeypatch.setattr(ps, "get_expert_model_parallel_group", lambda **_: expert) + monkeypatch.setattr(ps, "get_context_parallel_group", lambda **_: context) + monkeypatch.setattr(_impl.dist, "get_process_group_ranks", lambda g: ranks[id(g)]) + assert _impl._ep_group_is_cp_group(shape) + # EP spanning ranks outside this CP group sees other batches' rows. + ranks[id(context)] = [0, 2] + assert not _impl._ep_group_is_cp_group(shape) + monkeypatch.setattr(ps, "get_context_parallel_group", lambda **_: None) + assert not _impl._ep_group_is_cp_group(shape) + + +def test_routed_rows_use_the_cp_group_share_only_when_it_is_the_ep_group( + monkeypatch, +): + rank = _rank() + plan = rank._plan_flat_forward( + [ForwardInput(input_tokens=torch.arange(64), target_tokens=torch.arange(64))] + ) + assert rank._plan_group_routed_rows(plan) == (64,) + plan = replace(plan, signature=replace(plan.signature, topology=(1, 1, 2, 1))) + monkeypatch.setattr(rank, "_topology", lambda: SimpleNamespace(tp=1, cp=2)) + monkeypatch.setattr(rank, "_plan_group_rows", lambda plan: ((52480, True),)) + totals = [96794] + monkeypatch.setattr( + rank, "_cp_group_model_tokens", lambda batch, topology: totals[0] + ) + assert rank._plan_group_routed_rows(plan) == (52480,) + rank._ep_group_is_cp_group = True + assert rank._plan_group_routed_rows(plan) == (48397,) + totals[0] = 10**6 + assert rank._plan_group_routed_rows(plan) == (52480,) + + +def test_cp_group_total_is_its_larger_layout(monkeypatch): + runtime = pytest.importorskip("art.megatron.context_parallel.runtime") + bundle = SimpleNamespace( + token_layout_index=SimpleNamespace(token_counts_by_rank=(52480, 44314)) + ) + monkeypatch.setattr( + runtime, + "_get_or_build_planning_bundle", + lambda **_: ("key", bundle, None, None), + ) + monkeypatch.setattr( + runtime, + "_plan_gdn_global_execution", + lambda **_: SimpleNamespace(gdn_token_counts_by_rank=(48100, 48800)), + ) + values: dict[str, Any] = dict( + group_ids=None, parent_ids=None, topology=None, config=None, original_seq_len=0 + ) + total = runtime.context_parallel_model_token_total + assert total(**values, build_gdn_execution_spec=False) == 96794 + assert total(**values, build_gdn_execution_spec=True) == 96900 + # An empty rank still dispatches, and routes, one padding row. + bundle.token_layout_index.token_counts_by_rank = (2, 0) + assert total(**values, build_gdn_execution_spec=False) == 3 + + def test_split_charges_the_largest_hybridep_growth_beside_any_child_peak(): from art.trainer_rank._impl import _SubforwardCost @@ -700,7 +801,8 @@ def hybrid_checkpoint_rank(layer, monkeypatch): ), ) ) - assert rank._moe_output_bytes_per_token == 282624 + # This branch prices HybridEP routed rows on the EP group's balanced share. + assert rank._moe_output_bytes_per_token == 217908 assert rank._moe_memory_supported return rank @@ -742,7 +844,8 @@ def test_hybridep_recompute_prices_fresh_dense_output_without_buffer_growth( ) assert rank._plan_hybridep_growth_bytes(plan) == 0 retained, workspace = rank._checkpoint_memory_floor(groups) - assert workspace == 218752 * 2048 * 2 == 896008192 + # The combine output, with the TE workspaces live beside it. + assert workspace == 218752 * 2048 * 2 + rank._te_workspace_growth_bytes() cost = rank._subforward_cost(**values) # Unprofiled: the first execution's transients sit beside the extent. assert cost.required == int((8 + 2 * retained + workspace + COLD) * 1.1) diff --git a/tests/unit/test_trainer_rank_pending_memory.py b/tests/unit/test_trainer_rank_pending_memory.py index 8cae0619a..35f70c535 100644 --- a/tests/unit/test_trainer_rank_pending_memory.py +++ b/tests/unit/test_trainer_rank_pending_memory.py @@ -22,7 +22,7 @@ def module(cls): return obj -def rank_with_moe(moe_layer, *, install_hooks=False): +def rank_with_moe(moe_layer, *, install_hooks=False, stand_in=True): from megatron.core.ssm.gated_delta_net import GatedDeltaNet from megatron.core.transformer.transformer_block import TransformerBlock from transformer_engine.pytorch import RMSNorm @@ -97,6 +97,10 @@ def rank_with_moe(moe_layer, *, install_hooks=False): ) ) r._dp_rank_and_size = lambda: (0, 1) # Uninitialized MCore has no CPU DP group. + if stand_in: + # The one MoE layer stands in for all 40 of Qwen3.6-35B-A3B's: recompute + # is covered if that layer prices its FC1 stage too. + r._moe_recompute_covered = r._moe_gradient_enclosed == (True,) return r, gd @@ -134,7 +138,7 @@ def test_actual_constructor_cache_and_full_plan(pending_rank): assert ( rank._memory_check(plan).estimated_required_bytes == rank._plan_cost(plan).required - == 32303322409 + == 23404942633 ) selected = rank._select_next_micro_batch(requests, 0) assert ( @@ -197,7 +201,7 @@ def test_exact_pending_demand_survives_recovery(monkeypatch, pending_rank, fits_ plan = pending_rank._plan_flat_forward(requests) assert pending_rank._estimate_flat_forward(requests) is None assert g.plan_floor(pending_rank, plan) == (8296857600, 12705630112) - assert pending_rank._memory_check(plan).estimated_required_bytes == 32303322409 + assert pending_rank._memory_check(plan).estimated_required_bytes == 23404942633 _check_component_demand_recovery( monkeypatch, pending_rank, requests, fits_after=fits_after ) @@ -215,8 +219,8 @@ def test_original_installed_norm_preserves_pending_floor(layer): assert g.model_shapes(rank) is not None plan = rank._plan_flat_forward(full_requests()) assert g.plan_floor(rank, plan) == (8296857600, 12705630112) - assert rank._memory_check(plan).estimated_required_bytes == 32303322409 - assert rank._plan_cost(plan).required == 32303322409 + assert rank._memory_check(plan).estimated_required_bytes == 23404942633 + assert rank._plan_cost(plan).required == 23404942633 assert rank._estimate_flat_forward(full_requests()) is None for requests in ([], full_requests(no_grad=True)): assert g.plan_floor(rank, rank._plan_flat_forward(requests)) == (0, 0) @@ -383,9 +387,11 @@ def test_constructor_declined_moe_keeps_generic_admission(layer, unsupported): plan = rank._plan_flat_forward(requests) assert g.plan_floor(rank, plan) == (0, 0) required = rank._plan_cost(plan).required - # Generic checkpoint-input accounting still applies without a MoE component. + # Generic checkpoint accounting, including the recomputed layer's mixer, + # still applies without a MoE component. gradient = 50640 * 40 * 2048 * 2 - assert required == int((plan.output_bytes + 2 * gradient + COLD) * 1.1) + mixer = 50640 * rank._recomputed_mixer_bytes_per_token() + assert required == int((plan.output_bytes + 2 * gradient + mixer + COLD) * 1.1) rank._available_memory_bytes = lambda: required - 1 assert not rank._memory_check(plan).fits rank._available_memory_bytes = lambda: required diff --git a/tests/unit/test_trainer_rank_shared_memory.py b/tests/unit/test_trainer_rank_shared_memory.py index d8b1a2938..9016f6cc8 100644 --- a/tests/unit/test_trainer_rank_shared_memory.py +++ b/tests/unit/test_trainer_rank_shared_memory.py @@ -1,5 +1,6 @@ """One supported shared return held across routed compute; not all backward saves.""" +import math from types import SimpleNamespace import pytest @@ -106,6 +107,9 @@ def test_shared_return_in_actual_constructor_and_plan(layer, gate, no_grad): assert rank._moe_output_bytes_per_token == 192512 checkpoint_coefficient = 196608 if gate else 192512 assert rank._moe_checkpoint_grad_bytes_per_token == checkpoint_coefficient + # The shared return is the part that stays on local rows under HybridEP. + assert rank._moe_forward_shared_bytes == 4096 + assert rank._moe_gradient_shared_bytes == (8192 if gate else 4096) shapes = g.model_shapes(rank) assert shapes is not None and shapes[1][0].moe_bytes_per_row == 192512 requests = full_requests(no_grad) @@ -126,7 +130,8 @@ def test_shared_return_in_actual_constructor_and_plan(layer, gate, no_grad): 8296857600, 50640 * (checkpoint_coefficient + 128) + 3157761952, ) - assert rank._plan_cost(plan).required == (32759649577 if gate else 32531485993) + # The incoming gradient replaces one gradient per boundary. + assert rank._plan_cost(plan).required == (23861269801 if gate else 23633106217) selected = rank._select_next_micro_batch(requests, 0) assert ( selected.check.estimated_required_bytes @@ -147,7 +152,7 @@ def test_original_norm_installation_preserves_shared_return(layer, gated): 8296857600, 50640 * (checkpoint_coefficient + 128) + 3157761952, ) - expected = 32759649577 if gated else 32531485993 + expected = 23861269801 if gated else 23633106217 assert rank._memory_check(plan).estimated_required_bytes == expected assert rank._plan_cost(plan).required == expected @@ -305,14 +310,18 @@ def test_pre_gate_cache_precedes_owned_dispatcher_and_is_checkpoint_only(layer): ) # Installed dispatcher partials must not be repriced. assert rank._moe_checkpoint_grad_bytes_per_token == 196608 groups = ((19, True), (23, False)) + # Gradient rows also keep the recomputed layer's residual and norm rows and + # its MoE routing state; the first call allocates TE's cuBLAS workspaces. + beside = 2 * 2048 * 2 + rank._moe_checkpoint_state_bytes_per_token() assert rank._checkpoint_memory_floor(groups) == ( 19 * 40 * 4096, max( - 19 * (196608 + 128), + 19 * (196608 + 128 + beside), 23 * (192512 + 4 * 2048 * 2), - 19 * (196608 - 32768 + 128) + 10485760, + 19 * (196608 - 32768 + 128 + beside) + 10485760, 23 * (192512 - 32768 + 128 + 4 * 2048 * 2) + 10485760, - ), + ) + + rank._te_workspace_growth_bytes(), ) for mode in (None, "selective"): rank.runtime.model[0].decoder.config.recompute_granularity = mode @@ -340,7 +349,7 @@ def test_pre_gate_mixed_reference_and_exact_cost_mode_selection(layer, gradient_ ) assert rank._checkpoint_memory_floor(rank._plan_group_rows(mixed)) == ( 67 * 40 * 4096, - 4096 * (192512 + 4 * 2048 * 2), + 4096 * (192512 + 4 * 2048 * 2) + rank._te_workspace_growth_bytes(), ) # A reference-only path must not read or validate the unused gradient cache. rank._moe_checkpoint_grad_bytes_per_token = None @@ -410,11 +419,17 @@ def test_shared_return_escapes_the_ep_routed_allowance(layer, checkpoint_grad, s layer.config.context_parallel_size = layer.config.expert_model_parallel_size = 2 # EP2 halves the experts each rank owns. _hybridep(layer, 2).token_dispatcher.num_local_experts = local_experts // 2 - # HybridEP's 1.5x allowance turns top-k 8 into 12 routed rows; the gated - # shared return (doubled for checkpoint backward) is per local token. + # HybridEP's EP2 allowance (1.4) turns top-k 8 into 11.2 routed rows, each + # with one dispatched H-wide input; the gated shared return (doubled for + # checkpoint backward) is per local token. + collected: list[int] = [] assert ( _moe_output_bytes_per_token( - [layer], ParallelShape(tp=1, cp=2, ep=2), checkpoint_grad=checkpoint_grad + [layer], + ParallelShape(tp=1, cp=2, ep=2), + checkpoint_grad=checkpoint_grad, + shared_bytes=collected, ) - == 12 * 11776 * 2 + shared + == math.ceil(8 * 1.4 * 9728 * 2) + shared ) + assert collected == [shared] diff --git a/tests/unit/test_trainer_rank_slot_memory.py b/tests/unit/test_trainer_rank_slot_memory.py index 71ed04224..04bf8c06e 100644 --- a/tests/unit/test_trainer_rank_slot_memory.py +++ b/tests/unit/test_trainer_rank_slot_memory.py @@ -136,19 +136,27 @@ def test_subforward_cost_reuses_floor_without_caching_across_slot_changes( original = rank._checkpoint_memory_floor calls = [] - def floor(*args): - result = original(*args) + def floor(*args, **kwargs): + result = original(*args, **kwargs) calls.append(result) return result monkeypatch.setattr(rank, "_checkpoint_memory_floor", floor) + rows = rank._plan_group_rows(plan) + refs = tuple(group.slot_ref for group in plan.groups) + + def gradient(retained): + # Derived from the one floor call's boundaries, without another call. + return rank._checkpoint_input_gradient_bytes(rows, refs, retained=retained) + before = rank._plan_cost(plan) assert len(calls) == 1 - assert before.checkpoint_input_gradient == calls[0][0] + assert before.checkpoint_input_gradient == gradient(calls[0][0]) > 0 load_slot(rank, "selected", 64) after = rank._plan_cost(plan) assert len(calls) == 2 - assert after.checkpoint_input_gradient == calls[1][0] + assert after.checkpoint_input_gradient == gradient(calls[1][0]) > 0 + assert len(calls) == 2 assert calls[1][1] > calls[0][1] assert after.checkpoint_workspace > before.checkpoint_workspace assert after.required > before.required @@ -273,3 +281,19 @@ def guarded(name, *args, **kwargs): monkeypatch.setattr(builtins, "__import__", guarded) assert rank._slot_memory_shapes(ref) == () + + +def test_partial_slot_with_unpriced_fc1_stages_keeps_boundary_gradients(layer): + # A slot with FC1 adapters but no FC2 adapter still prices FC2 rows from the + # original metadata, but its FC1 converted weights go unpriced: keep one + # gradient per boundary for it. + rank, _ = rank_with_moe(weights(layer, 8)) + ref = load_slot(rank, "partial", 8) + groups = ((100, True),) + assert rank._moe_recompute_covered_for(ref) + assert rank._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 2048 * 2 + del layer.experts.linear_fc2.lora._slot_keys[ref] + assert rank._moe_workspace_bytes(1, checkpoint_grad=True, slot_ref=ref) > 0 + assert not rank._moe_recompute_covered_for(ref) + assert rank._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 40 * 4096 + assert rank._checkpoint_input_gradient_bytes(groups) == 100 * 2048 * 2 From e39b4603f5da326ff04e49eb279444ca4f131b75 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 05:50:15 +0000 Subject: [PATCH 38/40] Harden replay of the recomputed layer's readers Validate the recorded geometry and rank dimensions the recomputed-mixer arithmetic reads; accept any live Triton threshold and mirror live's empty head for non-positive rows; read the recompute and head-stage readers only with a checkpointed decoder, as live pricing does; check the producers of recorded staging and routed rows; and refuse a staged head that the recorded topology and head facts cannot produce. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_misses.py | 25 +++++- src/art/trainer_rank/_planner_replay.py | 52 +++++++++---- tests/unit/test_grouped_planner_replay.py | 93 ++++++++++++++++++++++- 3 files changed, 152 insertions(+), 18 deletions(-) diff --git a/src/art/trainer_rank/_planner_misses.py b/src/art/trainer_rank/_planner_misses.py index b1df930f1..f581ae9b4 100644 --- a/src/art/trainer_rank/_planner_misses.py +++ b/src/art/trainer_rank/_planner_misses.py @@ -15,7 +15,7 @@ from collections.abc import Callable, Mapping from contextlib import nullcontext -from dataclasses import asdict, dataclass +from dataclasses import MISSING, asdict, dataclass from datetime import datetime, timezone import hashlib from itertools import islice @@ -603,6 +603,15 @@ def replay( layers = values["num_layers"] if type(layers) is not int or not 0 < layers <= _planner_replay.MAX_LAYERS: raise ValueError("incomplete replay: recorded layer count out of bounds") + # Recomputed-layer pricing reads these alongside the geometry below. + if any( + type(values[name]) is not int or not 0 <= values[name] < 2**63 + for name in ("hidden_size", "param_dtype_size", "gdn_layers") + ) or any( + type(values[name]) is not bool + for name in ("sequence_parallel", "attention_output_gate") + ): + raise ValueError("incomplete replay: recorded rank dimensions are invalid") rank = _planner_replay.ReplayRank.__new__(_planner_replay.ReplayRank) for name in _RANK_FIELDS - {"one_layer_recompute"}: setattr(rank, "_" + name, values[name]) @@ -610,7 +619,19 @@ def replay( raise ValueError("incomplete replay: recompute mode is not recorded") rank._recorded_one_layer_recompute = values["one_layer_recompute"] rank._moe_forward_stages = tuple(tuple(row) for row in values["moe_forward_stages"]) - rank._geometry = ModelGeometry(**values["geometry"]) + # Recomputed-layer pricing multiplies these; accept only what live + # construction (ModelGeometry.from_config) records. Reports may omit + # fields that default to zero. + geometry = values["geometry"] + fields = ModelGeometry.__dataclass_fields__ + required = {name for name, field in fields.items() if field.default is MISSING} + if ( + type(geometry) is not dict + or not required <= set(geometry) <= set(fields) + or any(type(v) is not int or not 0 <= v < 2**63 for v in geometry.values()) + ): + raise ValueError("incomplete replay: recorded model geometry is invalid") + rank._geometry = ModelGeometry(**geometry) dp, tp, cp, pp = values["topology"] rank._topology_key = lambda: (dp, tp, cp, pp) # Check the same shared subforward inventories as capture, including the diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index 68890c1c3..2c9d9e747 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -54,6 +54,15 @@ ) +# Frozen per plan; None where no checkpointed decoder would read them. +_RECOMPUTE_READERS = ( + "moe_checkpoint_state_bytes_per_token", + "te_workspace_growth_bytes", + "backward_row_state_bytes", + "triton_min_rows", +) + + def refusal_reason(error: Exception) -> str: # Never retain arbitrary reader exception text, model reprs or token values. if ( @@ -128,6 +137,10 @@ def capture(rank: Any, plan: Any) -> dict[str, Any]: "_recomputed_mixer_bytes_per_token", "_mixer_activation_widths", "_triton_min_rows", + # Producers of recorded arguments replay checks against the facts. + "_plan_group_routed_rows", + "_plan_head_backward_traced", + "_head_backward_traced", ): method = getattr(rank, name) expected = getattr(_impl.TrainerRank, name) @@ -304,14 +317,12 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: "version": 3, "checkpoint_layers": layers, "checkpoint_moe_bytes_per_token": rank._checkpoint_moe_bytes_per_token(), - # Model and process readers of the recomputed layer and staged head. - "moe_checkpoint_state_bytes_per_token": ( - rank._moe_checkpoint_state_bytes_per_token() - ), - "te_workspace_growth_bytes": rank._te_workspace_growth_bytes(), - # Only a checkpointed Megatron decoder stages the head and has this. - "backward_row_state_bytes": rank._backward_row_state_bytes() if layers else 0, - "triton_min_rows": rank._triton_min_rows(), + # Model and process readers of the recomputed layer and staged head, + # read (like live pricing) only with a checkpointed decoder. + **{ + name: getattr(rank, "_" + name)() if layers else None + for name in _RECOMPUTE_READERS + }, "head_vocabulary": vocabulary, "head_target_backward": target_backward, "groups": groups, @@ -353,13 +364,17 @@ def integer(value: Any, *, minimum: int = 0) -> None: for key in ( "checkpoint_layers", "checkpoint_moe_bytes_per_token", - "moe_checkpoint_state_bytes_per_token", - "te_workspace_growth_bytes", - "backward_row_state_bytes", "head_vocabulary", ): integer(facts[key]) - integer(facts["triton_min_rows"], minimum=1) + for key in _RECOMPUTE_READERS: + if not facts["checkpoint_layers"]: + if facts[key] is not None: + raise ValueError("invalid runtime dimension") + continue + # Any integer Triton setting is live-accepted (a non-positive one + # prices no fallback chunk rows). + integer(facts[key], minimum=-(2**63) if key == "triton_min_rows" else 0) if type(facts["head_target_backward"]) is not bool: raise ValueError("invalid head backward eligibility") groups = facts["groups"] @@ -570,7 +585,8 @@ def _pending_adapter_gradient_bytes(self, refs: Any) -> tuple[int, ...]: def _head_workspace_bytes(self, rows: int) -> int: assert self._facts is not None - return _memory._dense_head_bytes(self._facts["head_vocabulary"], rows) + vocabulary = self._facts["head_vocabulary"] + return _memory._dense_head_bytes(vocabulary, rows) if rows > 0 else 0 def verify_group( self, group: dict[str, Any], layout: Any, records: list[dict[str, Any]] @@ -721,8 +737,16 @@ def runtime_arguments( recorded = arguments.get("group_routed_rows") if not isinstance(recorded, (list, tuple)) or tuple(recorded) != routed: raise ValueError("runtime facts disagree with selected routed rows") - if type(arguments.get("head_backward_traced")) is not bool: + traced = arguments.get("head_backward_traced") + if type(traced) is not bool: raise ValueError("head backward staging is not recorded") + # Live staging needs CP2 and a target backward through a priced head. + if traced and not ( + self._topology_key()[2] == 2 + and facts["head_target_backward"] + and facts["head_vocabulary"] + ): + raise ValueError("head backward staging disagrees with recorded facts") head = max( max( _memory._dense_head_bytes(facts["head_vocabulary"], g["head_rows"]), diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 23a93a29a..46ea43346 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -329,7 +329,7 @@ def test_recomputed_layer_fact_validation_rejects_forged_input(change, layer, tm facts["groups"][0]["moe_covered"] = True message = "invalid MoE recompute coverage" elif change == "triton_min_rows": - facts["triton_min_rows"] = 0 + facts["triton_min_rows"] = 64.0 message = "invalid runtime dimension" elif change == "routed_rows": item["arguments"]["group_routed_rows"][1] -= 1 @@ -341,9 +341,98 @@ def test_recomputed_layer_fact_validation_rejects_forged_input(change, layer, tm reports.replay(report) +@pytest.mark.parametrize("value", [None, 1.5, -1]) +def test_replay_refuses_invalid_recorded_geometry(value, tmp_path): + rank = head_rank() + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, _ = emitted( + rank, rank._plan_flat_forward([request(65, grad=True)]), tmp_path + ) + # The recomputed mixer multiplies these; refuse before any pricing. + report["replay"]["memory_replay"]["rank"]["geometry"]["kv_channels"] = value + with pytest.raises(ValueError, match="geometry"): + reports.replay(report) + + +@pytest.mark.parametrize("setting", ["0", "-3"]) +def test_non_positive_triton_threshold_is_captured_and_replayed( + setting, monkeypatch, tmp_path +): + # Live pricing accepts it (the fallback chunk then prices no rows). + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", setting) + rank = head_rank() + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, costs = emitted( + rank, rank._plan_flat_forward([request(65, grad=True)]), tmp_path + ) + facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] + assert facts["triton_min_rows"] == int(setting) + result = reports.replay(report) + assert result["aggregate"]["matches"] + assert result["estimates"][0]["required_bytes"] == costs[0].required + + +def test_recompute_readers_are_unread_without_a_checkpointed_decoder( + monkeypatch, tmp_path +): + from art.trainer_rank import _planner_replay + + # Live pricing reads them only for a checkpointed decoder's recompute; + # capture must not refuse a plan over a setting live never reads. + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", "not-a-number") + rank = _rank(monkeypatch) + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, _ = emitted(rank, rank._plan_flat_forward([_request(1)]), tmp_path) + facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] + assert facts["checkpoint_layers"] == 0 + assert all(facts[name] is None for name in _planner_replay._RECOMPUTE_READERS) + assert reports.replay(report)["aggregate"]["matches"] + forged = deepcopy(report) + forged["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "te_workspace_growth_bytes" + ] = 0 + with pytest.raises(ValueError, match="invalid runtime dimension"): + reports.replay(forged) + + +def test_staged_head_needs_recorded_cp2_target_backward(tmp_path): + rank = head_rank() + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, _ = emitted( + rank, rank._plan_flat_forward([request(65, grad=True)]), tmp_path + ) + # This CP1 plan cannot stage its head; a forged flag must not price it so. + report["replay"]["memory_replay"]["estimates"][0]["arguments"][ + "head_backward_traced" + ] = True + with pytest.raises(ValueError, match="staging disagrees"): + reports.replay(report) + + +@pytest.mark.parametrize( + "field,value", [("hidden_size", "8"), ("gdn_layers", -1), ("sequence_parallel", 0)] +) +def test_replay_refuses_invalid_recorded_rank_dimensions(field, value, tmp_path): + rank = head_rank() + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, _ = emitted( + rank, rank._plan_flat_forward([request(65, grad=True)]), tmp_path + ) + report["replay"]["memory_replay"]["rank"][field] = value + with pytest.raises(ValueError, match="rank dimensions"): + reports.replay(report) + + @pytest.mark.parametrize( "name", - ["_te_workspace_growth_bytes", "_moe_recompute_covered_for", "_triton_min_rows"], + [ + "_te_workspace_growth_bytes", + "_moe_recompute_covered_for", + "_triton_min_rows", + "_plan_group_routed_rows", + "_plan_head_backward_traced", + "_head_backward_traced", + ], ) def test_custom_recomputed_layer_reader_is_explicitly_incomplete(name, monkeypatch): from art.trainer_rank import _planner_replay From bab7be10a856c386f56b31a3b920376bdf97dd87 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 06:21:41 +0000 Subject: [PATCH 39/40] Cover the staged CP2 head and recorded geometry in replay tests Update the fact-budget test's forged MoE terms to the shared-coefficient shape. Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_grouped_planner_replay.py | 60 +++++++++++++++++++ .../unit/test_planner_runtime_fact_guards.py | 2 +- 2 files changed, 61 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 46ea43346..d3540efe4 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -409,6 +409,66 @@ def test_staged_head_needs_recorded_cp2_target_backward(tmp_path): reports.replay(report) +def test_staged_cp2_head_is_replayed(monkeypatch, tmp_path): + monkeypatch.setattr( + tr, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, + ) + from types import SimpleNamespace + + from art.megatron.context_parallel.types import ParallelTopology + + # CP2 through the stock topology reader (as capture requires), over a CPU + # CP2 topology and planning config. + monkeypatch.setattr(tr.TrainerRank, "_topology_key", lambda self: (1, 1, 2, 1)) + rank = head_rank() + rank._topology = lambda: ParallelTopology(tp=1, cp=2) + for name, value in { + "linear_num_key_heads": 16, + "linear_num_value_heads": 32, + "linear_key_head_dim": 128, + "linear_value_head_dim": 128, + "params_dtype": torch.bfloat16, + }.items(): + setattr(rank.runtime.provider, name, value) + rank.runtime.model_support_handler = SimpleNamespace( + build_gdn_execution_spec=True, + context_parallel_workload_profile=lambda provider: None, + ) + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + plan = rank._plan_flat_forward([request(512, grad=True)]) + assert plan.signature.topology[2] == 2 + report, costs = emitted(rank, plan, tmp_path) + item = report["replay"]["memory_replay"]["estimates"][0] + assert item["arguments"]["head_backward_traced"] is True + result = reports.replay(report) + assert result["aggregate"]["matches"] + assert result["estimates"][0]["required_bytes"] == costs[0].required + + +@pytest.mark.parametrize("change", ["omit_default", "omit_required", "extra"]) +def test_recorded_geometry_fields_follow_its_schema(change, tmp_path): + rank = head_rank() + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, _ = emitted( + rank, rank._plan_flat_forward([request(65, grad=True)]), tmp_path + ) + geometry = report["replay"]["memory_replay"]["rank"]["geometry"] + if change == "omit_default": + # Older reports may omit fields that default to zero. + (name,) = [n for n, v in geometry.items() if n == "gdn_conv_kernel" and v == 0] + del geometry[name] + assert reports.replay(report)["aggregate"]["matches"] + return + if change == "omit_required": + del geometry["kv_channels"] + else: + geometry["unknown_width"] = 1 + with pytest.raises(ValueError, match="geometry"): + reports.replay(report) + + @pytest.mark.parametrize( "field,value", [("hidden_size", "8"), ("gdn_layers", -1), ("sequence_parallel", 0)] ) diff --git a/tests/unit/test_planner_runtime_fact_guards.py b/tests/unit/test_planner_runtime_fact_guards.py index aa2002960..9b2a6265f 100644 --- a/tests/unit/test_planner_runtime_fact_guards.py +++ b/tests/unit/test_planner_runtime_fact_guards.py @@ -56,7 +56,7 @@ def test_cumulative_fact_budget_precedes_json_encoding(inventory, monkeypatch): group = facts["groups"][0] group["gdn"] = None if inventory == "stages": - group["forward"] = [0, [[0, 0]] * 4096] + group["forward"] = [0, [[0, 0]] * 4096, 0] group["gradient"] = deepcopy(group["forward"]) facts["groups"] = [deepcopy(group) for _ in range(8)] elif inventory == "slots": From 1763c67dc2de42ef79e67ca97fbf8d11ea1e172c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 06:46:19 +0000 Subject: [PATCH 40/40] Run the head-stage memory tests in the Megatron lane Co-Authored-By: Claude Opus 5.5 (1M context) --- .github/workflows/prek.yml | 2 ++ 1 file changed, 2 insertions(+) diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index f7360485f..cce655bd8 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -237,6 +237,7 @@ jobs: tests/unit/test_trainer_rank_slot_memory.py \ tests/unit/test_trainer_rank_moe_memory.py \ tests/unit/test_trainer_rank_head_memory.py \ + tests/unit/test_trainer_rank_head_stage_memory.py \ tests/unit/test_grouped_planner_replay.py \ tests/unit/test_planner_runtime_fact_guards.py \ tests/unit/test_planner_replay_owner_budget.py \ @@ -287,6 +288,7 @@ jobs: --ignore=tests/unit/test_trainer_rank_slot_memory.py \ --ignore=tests/unit/test_trainer_rank_moe_memory.py \ --ignore=tests/unit/test_trainer_rank_head_memory.py \ + --ignore=tests/unit/test_trainer_rank_head_stage_memory.py \ --ignore=tests/unit/test_grouped_planner_replay.py \ --ignore=tests/unit/test_planner_runtime_fact_guards.py \ --ignore=tests/unit/test_planner_replay_owner_budget.py \