From 65076daac36415be36d373ffd0a84051cbde3ae3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 00:44:21 +0000 Subject: [PATCH 1/4] Price HybridEP expert work and buffer growth at EP>1 At EP>1 the routed-expert working set was priced at zero, so cold CP2/EP2 plans fell back to retention floors alone (8.4 GB predicted against 34 GB observed on 8 layers). HybridEP hands each rank the pairs routed to its local experts, so reuse the EP1 per-token coefficient with a 1.5x routing-imbalance allowance (a pretrained CP2/EP2 run put 1.35x the balanced load on one rank). HybridEP's persistent buffers are allocated outside the PyTorch allocator after admission, so charge any growth a plan triggers: capacity x EP x (2H + 5E + 4H/128) bytes, less what is already held. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 127 ++++++++++++++++++--- tests/unit/test_trainer_rank_moe_memory.py | 81 +++++++++++++ 2 files changed, 191 insertions(+), 17 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 5f0006c4a..fea8d3a6f 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1494,6 +1494,39 @@ 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 + + +def _moe_dispatcher_supported( + dispatcher: Any, ep: int, *, alltoall: type, flex: type, hybridep: type +) -> bool: + """EP1 all-to-all, or ART's HybridEP flex dispatcher across the EP group.""" + if ep == 1: + return type(dispatcher) is alltoall + return ( + type(dispatcher) is flex + and type(getattr(dispatcher, "_comm_manager", None)) is hybridep + and getattr(dispatcher, "ep_size", None) == ep + and getattr(dispatcher, "tp_size", None) == 1 + ) + + +def _hybridep_buffer_bytes(capacity: int, ranks: int, hidden: int, experts: int) -> int: + """Intranode HybridEP buffers for a per-rank token capacity. + + Dispatch outputs alias the combine inputs (the shared-buffer default), sized + for every rank's tokens routed to one rank: BF16 tokens, FP32 probabilities + over the node's experts and FP32 FP8 scaling factors, which are allocated + even without FP8. The routing-map allgather keeps one byte per expert. + """ + if capacity <= 0: + return 0 + tokens = capacity * ranks + tokens += -tokens % 4 + return tokens * (2 * hidden + 5 * experts + 4 * (hidden // 128)) + + def _moe_output_bytes_per_token( model: Sequence[torch.nn.Module], shape: ParallelShape, @@ -1503,8 +1536,9 @@ def _moe_output_bytes_per_token( slot_ref: "LoRASlotRef | None" = None, ) -> int: """Known routed-expert working set, not a complete model/compiled bound.""" - # CP shards rows, not the per-token working set; EP dispatch is not modeled. - if (shape.tp, shape.ep, shape.etp) != (1, 1, 1): + # 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): return 0 from megatron.core.extensions.transformer_engine import ( TEColumnParallelGroupedLinear, @@ -1515,10 +1549,16 @@ def _moe_output_bytes_per_token( from megatron.core.transformer.moe.router import TopKRouter from megatron.core.transformer.moe.token_dispatcher import ( MoEAlltoAllTokenDispatcher, + MoEFlexTokenDispatcher, + _HybridEPManager, ) 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 = min(shape.ep, _EP_ROUTED_ROW_ALLOWANCE) if shape.ep > 1 else 1 coefficient = 0 for chunk in model: for layer in chunk.modules(): @@ -1536,15 +1576,18 @@ def _moe_output_bytes_per_token( (getattr(fc2, "linear_fc2", None), TERowParallelGroupedLinear), (getattr(layer, "router", None), TopKRouter), ) - if ( - any( - type(site) is not expected - or "forward" in vars(site) - or getattr(site, "_forward_hooks", None) - or getattr(site, "_forward_pre_hooks", None) - for site, expected in sites - ) - or type(dispatcher) is not MoEAlltoAllTokenDispatcher + if any( + type(site) is not expected + or "forward" in vars(site) + or getattr(site, "_forward_hooks", None) + or getattr(site, "_forward_pre_hooks", None) + for site, expected in sites + ) or not _moe_dispatcher_supported( + dispatcher, + shape.ep, + alltoall=MoEAlltoAllTokenDispatcher, + flex=MoEFlexTokenDispatcher, + hybridep=_HybridEPManager, ): return 0 config = layer.config @@ -1609,7 +1652,9 @@ def _moe_output_bytes_per_token( and fc1.fused_gate_up and not fc1.non_gated and fc1.out_features == 2 * inputs.shape[-2] - and getattr(dispatcher, "ep_size", None) == 1 + # 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 and getattr(config, "moe_permute_fusion", False) @@ -1632,15 +1677,14 @@ 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 - row_bytes = ( - config.moe_router_topk * features * weights.element_size() + shared - ) + routed_rows = math.ceil(config.moe_router_topk * routed_allowance) + row_bytes = 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: padded, transposes, effective = storage saved_fc1, rank_fc1 = 0, 0 - routed_size = config.moe_router_topk * weights.element_size() + routed_size = routed_rows * weights.element_size() if enclosing_fc1 is not None: adapter = getattr(enclosing_fc1, "lora", None) base = getattr(enclosing_fc1, "linear_fc1", None) @@ -3794,6 +3838,45 @@ def _plan_head_workspace_bytes(self, plan: _FlatForwardPlan) -> int: ) return peak + def _plan_hybridep_growth_bytes(self, plan: _FlatForwardPlan) -> int: + """HybridEP buffer growth this plan triggers before its forward. + + The buffer is allocated outside the PyTorch allocator after admission, + so neither the free-memory sample nor a learned peak includes it. + """ + provider: Any = getattr(getattr(self, "runtime", None), "provider", None) + ep = int(getattr(provider, "expert_model_parallel_size", 1) or 1) + if ep <= 1 or not plan.groups: + return 0 + from megatron.core.transformer.moe import fused_a2a + + from art.megatron.train import _hybridep_token_capacity + + topology = self._topology() + sequence = max( + int( + _pad_packed_batch(group.packed, multiple=int(topology.tp)).tokens.shape[ + 1 + ] + ) + for group in plan.groups + ) + rows = max(rows for rows, _ in self._plan_group_rows(plan)) + capacity = max(_hybridep_token_capacity(sequence, int(topology.cp)), rows) + current = fused_a2a._hybrid_ep_buffer + held = ( + 0 + if current is None + else int(current.configurer.buffer_config.max_num_of_tokens_per_rank) + ) + if capacity <= held: + return 0 + hidden = int(provider.hidden_size) + experts = int(provider.num_moe_experts) + return _hybridep_buffer_bytes(capacity, ep, hidden, experts) - ( + _hybridep_buffer_bytes(held, ep, hidden, experts) + ) + def _plan_group_rows(self, plan: _FlatForwardPlan) -> tuple[tuple[int, bool], ...]: """Physical rows per group on the most loaded context-parallel rank.""" topology = self._topology() if plan.signature.topology[2] > 1 else None @@ -3972,6 +4055,7 @@ def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: head_workspace_bytes=self._plan_head_workspace_bytes(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), ) def _subforward_cost( @@ -3987,6 +4071,7 @@ def _subforward_cost( head_workspace_bytes: int = 0, checkpoint_floor: tuple[int, int] = (0, 0), retained_tokens: int | None = None, + hybridep_growth_bytes: int = 0, ) -> _SubforwardCost: required = self._estimate_required_memory_bytes_from_values( packed_tokens=packed_tokens, @@ -4000,6 +4085,7 @@ def _subforward_cost( checkpoint_floor=checkpoint_floor, retained_tokens=retained_tokens, include_checkpoint_input_gradient=False, + hybridep_growth_bytes=hybridep_growth_bytes, ) checkpoint_retained, checkpoint_workspace = self._checkpoint_memory_floor( group_rows, slot_refs @@ -6097,6 +6183,9 @@ def _fill_planner_snapshot( "gdn_segments": child.grad_segment_count, "retained_tokens": self._plan_retained_tokens(child), "group_rows": self._plan_group_rows(child), + "hybridep_growth_bytes": ( + self._plan_hybridep_growth_bytes(child) + ), }, "expected_required_bytes": cost.required, "retained_bytes": cost.retained, @@ -6822,6 +6911,7 @@ def _memory_check( head_workspace_bytes=self._plan_head_workspace_bytes(forward), checkpoint_floor=_gdn_memory.plan_floor(self, forward), retained_tokens=self._plan_retained_tokens(forward), + hybridep_growth_bytes=self._plan_hybridep_growth_bytes(forward), ) return self._memory_check_required(required, sync_across_dp=sync_across_dp) @@ -7669,6 +7759,7 @@ def _estimate_required_memory_bytes_from_values( checkpoint_floor: tuple[int, int] = (0, 0), retained_tokens: int | None = None, include_checkpoint_input_gradient: bool = True, + hybridep_growth_bytes: int = 0, ) -> int: if packed_tokens <= 0: return output_bytes @@ -7823,7 +7914,9 @@ def _estimate_required_memory_bytes_from_values( ) ), ) - return int((output_bytes + compute) * _MEMORY_SAFETY_FACTOR) + return int( + (output_bytes + compute + hybridep_growth_bytes) * _MEMORY_SAFETY_FACTOR + ) def _one_layer_recompute(self) -> bool: """ART's default full/uniform/1 recompute, which Megatron runs in training. diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 33b8ce9df..4ed36c6dc 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -188,6 +188,41 @@ def test_sharded_path_unchanged(layer, field): ) +def _hybridep(layer, ep: int, manager: str = "hybridep"): + from megatron.core.transformer.moe import token_dispatcher + + dispatcher = object.__new__(token_dispatcher.MoEFlexTokenDispatcher) + dispatcher.config = layer.config + dispatcher.ep_size, dispatcher.tp_size = ep, 1 + managers = { + "hybridep": token_dispatcher._HybridEPManager, + "deepep": token_dispatcher._DeepepManager, + } + dispatcher._comm_manager = object.__new__(managers[manager]) + layer.token_dispatcher = dispatcher + return layer + + +@pytest.mark.parametrize("ep", [2, 4]) +def test_hybridep_prices_routed_rows_with_imbalance_allowance(layer, ep): + # 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. + 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 + + +@pytest.mark.parametrize( + "manager,ep,shape_ep", [("deepep", 2, 2), ("hybridep", 2, 4), ("hybridep", 1, 1)] +) +def test_other_flex_dispatchers_stay_unmodeled(layer, manager, ep, shape_ep): + dispatcher = _hybridep(layer, ep, manager) + shape = ParallelShape(tp=1, cp=1, ep=shape_ep) + assert _moe_output_bytes_per_token([dispatcher], shape) == 0 + + @pytest.mark.parametrize( "field,value", [ @@ -492,3 +527,49 @@ def test_unknown_enclosing_lifetimes_keep_previous_fc2_floor( } setattr(sites[site], attribute, value) assert _moe_output_bytes_per_token([layer], ParallelShape(tp=1, cp=1)) == 106_496 + + +def test_hybridep_buffer_growth_is_charged_before_forward(monkeypatch): + from megatron.core.transformer.moe import fused_a2a + + from art.trainer_rank._impl import _hybridep_buffer_bytes + + # Two ranks' tokens to one rank: BF16 H, FP32 probs and a routing-map byte + # per expert (256), and FP32 scaling factors per H/128. + assert _hybridep_buffer_bytes(1000, 2, 2048, 256) == 2000 * (4096 + 1280 + 64) + assert _hybridep_buffer_bytes(0, 2, 2048, 256) == 0 + rank = _rank() + rank.runtime.provider.expert_model_parallel_size = 2 + rank.runtime.provider.num_moe_experts = 256 + plan = rank._plan_flat_forward( + [ForwardInput(input_tokens=torch.arange(64), target_tokens=torch.arange(64))] + ) + monkeypatch.setattr(rank, "_topology", lambda: SimpleNamespace(tp=1, cp=2)) + monkeypatch.setattr(rank, "_plan_group_rows", lambda plan: ((600, True),)) + monkeypatch.setattr( + "art.megatron.train._hybridep_token_capacity", lambda sequence, cp: 1000 + ) + + def held(capacity): + config = SimpleNamespace(max_num_of_tokens_per_rank=capacity) + return SimpleNamespace(configurer=SimpleNamespace(buffer_config=config)) + + full = _hybridep_buffer_bytes(1000, 2, 2048, 256) + for current, growth in ( + (None, full), + (held(400), full - _hybridep_buffer_bytes(400, 2, 2048, 256)), + (held(1000), 0), + ): + monkeypatch.setattr(fused_a2a, "_hybrid_ep_buffer", current) + assert rank._plan_hybridep_growth_bytes(plan) == growth + # The growth enters required memory, not forward retention. + rank._update_memory_profile(plan, 10**9, retained_bytes=10**8) + monkeypatch.setattr(fused_a2a, "_hybrid_ep_buffer", None) + grown = rank._plan_cost(plan) + monkeypatch.setattr(rank, "_plan_hybridep_growth_bytes", lambda plan: 0) + base = rank._plan_cost(plan) + assert grown.required - base.required == pytest.approx(full * 1.1, abs=2) + assert grown.retained == base.retained < base.required + rank.runtime.provider.expert_model_parallel_size = 1 + monkeypatch.undo() + assert rank._plan_hybridep_growth_bytes(plan) == 0 From d55d96908cfbe99027d0a511f70a69d8fa8214c0 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 01:03:59 +0000 Subject: [PATCH 2/4] Charge HybridEP growth once, at full size, in both peaks Review follow-ups: the growth now enters the forward peak and, via the max-combined checkpoint workspace, the checkpoint peak, but not retention, so split plans pay it once. The replacement buffer is charged in full because the old one stays referenced while it is allocated. Rows per rank follow HybridEP's TMA alignment, 512-row minimum and 64-row chunks, over the ETPxEP group and its expert columns. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 54 +++++++++++------- tests/unit/test_trainer_rank_moe_memory.py | 66 ++++++++++++++++------ 2 files changed, 82 insertions(+), 38 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index fea8d3a6f..e6a0db998 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1512,18 +1512,27 @@ def _moe_dispatcher_supported( ) +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.""" + multiple = 4 // math.gcd(4, ranks) + rows = max(-(-capacity // multiple) * multiple, 512) + return -(-rows // 64) * 64 + + def _hybridep_buffer_bytes(capacity: int, ranks: int, hidden: int, experts: int) -> int: """Intranode HybridEP buffers for a per-rank token capacity. - Dispatch outputs alias the combine inputs (the shared-buffer default), sized - for every rank's tokens routed to one rank: BF16 tokens, FP32 probabilities - over the node's experts and FP32 FP8 scaling factors, which are allocated - even without FP8. The routing-map allgather keeps one byte per expert. + ``ranks`` is the ETPxEP communication group and ``experts`` its expert + columns. Dispatch outputs alias the combine inputs (the shared-buffer + default), sized for every rank's tokens routed to one rank: BF16 tokens, + FP32 probabilities over the columns and FP32 FP8 scaling factors, which are + allocated even without FP8. The routing-map allgather keeps one byte per + column. """ if capacity <= 0: return 0 - tokens = capacity * ranks - tokens += -tokens % 4 + tokens = _hybridep_rows_per_rank(capacity, ranks) * ranks return tokens * (2 * hidden + 5 * experts + 4 * (hidden // 128)) @@ -1558,7 +1567,7 @@ def _moe_output_bytes_per_token( # 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 = min(shape.ep, _EP_ROUTED_ROW_ALLOWANCE) if shape.ep > 1 else 1 + routed_allowance = _EP_ROUTED_ROW_ALLOWANCE if shape.ep > 1 else 1 coefficient = 0 for chunk in model: for layer in chunk.modules(): @@ -3869,12 +3878,16 @@ def _plan_hybridep_growth_bytes(self, plan: _FlatForwardPlan) -> int: if current is None else int(current.configurer.buffer_config.max_num_of_tokens_per_rank) ) - if capacity <= held: + etp = int(getattr(provider, "expert_tensor_parallel_size", 1) or 1) + ranks = ep * etp + if _hybridep_rows_per_rank(capacity, ranks) <= held: return 0 - hidden = int(provider.hidden_size) - experts = int(provider.num_moe_experts) - return _hybridep_buffer_bytes(capacity, ep, hidden, experts) - ( - _hybridep_buffer_bytes(held, ep, hidden, experts) + # The old buffer stays referenced while its replacement is allocated. + return _hybridep_buffer_bytes( + capacity, + ranks, + int(provider.hidden_size), + int(provider.num_moe_experts) * etp, ) def _plan_group_rows(self, plan: _FlatForwardPlan) -> tuple[tuple[int, bool], ...]: @@ -4085,7 +4098,6 @@ def _subforward_cost( checkpoint_floor=checkpoint_floor, retained_tokens=retained_tokens, include_checkpoint_input_gradient=False, - hybridep_growth_bytes=hybridep_growth_bytes, ) checkpoint_retained, checkpoint_workspace = self._checkpoint_memory_floor( group_rows, slot_refs @@ -4109,9 +4121,13 @@ def _subforward_cost( checkpoint_retained = output_bytes + max( checkpoint_retained, checkpoint_floor[0] ) - checkpoint_workspace = max( - checkpoint_workspace, head_workspace_bytes, checkpoint_floor[1] + # HybridEP buffer growth stays allocated through backward, but is not + # forward retention: split plans charge it once, via max workspace. + checkpoint_workspace = ( + max(checkpoint_workspace, head_workspace_bytes, checkpoint_floor[1]) + + hybridep_growth_bytes ) + required += int(hybridep_growth_bytes * _MEMORY_SAFETY_FACTOR) forward_required = required if gradient: required = max( @@ -6911,8 +6927,7 @@ def _memory_check( head_workspace_bytes=self._plan_head_workspace_bytes(forward), checkpoint_floor=_gdn_memory.plan_floor(self, forward), retained_tokens=self._plan_retained_tokens(forward), - hybridep_growth_bytes=self._plan_hybridep_growth_bytes(forward), - ) + ) + int(self._plan_hybridep_growth_bytes(forward) * _MEMORY_SAFETY_FACTOR) return self._memory_check_required(required, sync_across_dp=sync_across_dp) def _admission_outcome(self, local: int) -> int: @@ -7759,7 +7774,6 @@ def _estimate_required_memory_bytes_from_values( checkpoint_floor: tuple[int, int] = (0, 0), retained_tokens: int | None = None, include_checkpoint_input_gradient: bool = True, - hybridep_growth_bytes: int = 0, ) -> int: if packed_tokens <= 0: return output_bytes @@ -7914,9 +7928,7 @@ def _estimate_required_memory_bytes_from_values( ) ), ) - return int( - (output_bytes + compute + hybridep_growth_bytes) * _MEMORY_SAFETY_FACTOR - ) + return int((output_bytes + compute) * _MEMORY_SAFETY_FACTOR) def _one_layer_recompute(self) -> bool: """ART's default full/uniform/1 recompute, which Megatron runs in training. diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 4ed36c6dc..56b17af34 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -214,6 +214,19 @@ def test_hybridep_prices_routed_rows_with_imbalance_allowance(layer, ep): assert single > 0 and expert == single // 8 * 12 +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. + single = _moe_output_bytes_per_token( + [_enclosing_moe(layer)], ParallelShape(tp=1, cp=1) + ) + expert = _hybridep(layer, 2) + 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 + + @pytest.mark.parametrize( "manager,ep,shape_ep", [("deepep", 2, 2), ("hybridep", 2, 4), ("hybridep", 1, 1)] ) @@ -530,13 +543,19 @@ def test_unknown_enclosing_lifetimes_keep_previous_fc2_floor( def test_hybridep_buffer_growth_is_charged_before_forward(monkeypatch): - from megatron.core.transformer.moe import fused_a2a - - from art.trainer_rank._impl import _hybridep_buffer_bytes - - # Two ranks' tokens to one rank: BF16 H, FP32 probs and a routing-map byte - # per expert (256), and FP32 scaling factors per H/128. - assert _hybridep_buffer_bytes(1000, 2, 2048, 256) == 2000 * (4096 + 1280 + 64) + fused_a2a = pytest.importorskip("megatron.core.transformer.moe.fused_a2a") + from art.trainer_rank._impl import _hybridep_buffer_bytes, _hybridep_rows_per_rank + + # Rows per rank are TMA-aligned, at least 512 and padded to 64-row chunks; + # every rank's rows land on one rank: BF16 H, FP32 probs and a routing-map + # byte per expert column (256), and FP32 scaling factors per H/128. + assert [_hybridep_rows_per_rank(n, 2) for n in (1, 512, 1000, 1025)] == [ + 512, + 512, + 1024, + 1088, + ] + assert _hybridep_buffer_bytes(1000, 2, 2048, 256) == 2 * 1024 * (4096 + 1280 + 64) assert _hybridep_buffer_bytes(0, 2, 2048, 256) == 0 rank = _rank() rank.runtime.provider.expert_model_parallel_size = 2 @@ -550,26 +569,39 @@ def test_hybridep_buffer_growth_is_charged_before_forward(monkeypatch): "art.megatron.train._hybridep_token_capacity", lambda sequence, cp: 1000 ) - def held(capacity): - config = SimpleNamespace(max_num_of_tokens_per_rank=capacity) + def held(rows): + config = SimpleNamespace(max_num_of_tokens_per_rank=rows) return SimpleNamespace(configurer=SimpleNamespace(buffer_config=config)) + # The old buffer stays referenced while its replacement is allocated, so a + # growing plan pays the full new size. full = _hybridep_buffer_bytes(1000, 2, 2048, 256) - for current, growth in ( - (None, full), - (held(400), full - _hybridep_buffer_bytes(400, 2, 2048, 256)), - (held(1000), 0), - ): + for current, growth in ((None, full), (held(512), full), (held(1024), 0)): monkeypatch.setattr(fused_a2a, "_hybrid_ep_buffer", current) assert rank._plan_hybridep_growth_bytes(plan) == growth - # The growth enters required memory, not forward retention. - rank._update_memory_profile(plan, 10**9, retained_bytes=10**8) + # ETP multiplies the communication ranks and expert columns. monkeypatch.setattr(fused_a2a, "_hybrid_ep_buffer", None) + rank.runtime.provider.expert_tensor_parallel_size = 2 + assert rank._plan_hybridep_growth_bytes(plan) == _hybridep_buffer_bytes( + 1000, 4, 2048, 512 + ) + rank.runtime.provider.expert_tensor_parallel_size = 1 + # 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) + ) grown = rank._plan_cost(plan) monkeypatch.setattr(rank, "_plan_hybridep_growth_bytes", lambda plan: 0) base = rank._plan_cost(plan) - assert grown.required - base.required == pytest.approx(full * 1.1, abs=2) + charged = int(full * 1.1) assert grown.retained == base.retained < base.required + assert grown.checkpoint_workspace - base.checkpoint_workspace == full + assert grown.required - base.required in (charged, charged + 1) + # A split pays it once, however many children would grow the buffer. + split = rank._split_required_memory([grown, grown, grown]) + extra = split - rank._split_required_memory([base, base, base]) + assert charged - 1 <= extra <= charged + 1 rank.runtime.provider.expert_model_parallel_size = 1 monkeypatch.undo() assert rank._plan_hybridep_growth_bytes(plan) == 0 From 52dee283c037ec069e6789a3625244057383ce7b Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 01:12:41 +0000 Subject: [PATCH 3/4] Charge the largest HybridEP growth once across a split A grown buffer persists into later split children, so a child with the largest workspace can run beside another child's growth. Track growth as its own cost field, exclude it from each child's ephemeral and workspace terms, and add the largest growth once to both split aggregates. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 40 +++++++++++++++------- tests/unit/test_trainer_rank_moe_memory.py | 27 +++++++++++++-- 2 files changed, 52 insertions(+), 15 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index e6a0db998..d4718cab0 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1048,6 +1048,9 @@ class _SubforwardCost: checkpoint_input_gradient: int = 0 # Already in required; backward allowance must not reorder forward execution. checkpoint_peak_increment: int = 0 + # 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 @property def ephemeral(self) -> int: @@ -3340,16 +3343,28 @@ def _split_rung_check(self, costs: Sequence[_SubforwardCost]) -> _MemoryCheck: @staticmethod def _split_required_memory(costs: Sequence[_SubforwardCost]) -> int: - required = sum(cost.retained for cost in costs) + max( - cost.ephemeral for cost in costs + # A grown HybridEP buffer persists into later children: charge the + # largest growth once, beside whichever child peaks highest. + growth = max(cost.hybridep_growth for cost in costs) + required = ( + sum(cost.retained for cost in costs) + + max( + cost.ephemeral - int(cost.hybridep_growth * _MEMORY_SAFETY_FACTOR) + for cost in costs + ) + + int(growth * _MEMORY_SAFETY_FACTOR) ) 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. - checkpoint = sum( - cost.checkpoint_retained + cost.checkpoint_input_gradient - for cost in costs - ) + max(cost.checkpoint_workspace for cost in costs) + checkpoint = ( + sum( + cost.checkpoint_retained + cost.checkpoint_input_gradient + for cost in costs + ) + + max(cost.checkpoint_workspace for cost in costs) + + growth + ) required = max(required, int(checkpoint * _MEMORY_SAFETY_FACTOR)) return required @@ -4121,13 +4136,9 @@ def _subforward_cost( checkpoint_retained = output_bytes + max( checkpoint_retained, checkpoint_floor[0] ) - # HybridEP buffer growth stays allocated through backward, but is not - # forward retention: split plans charge it once, via max workspace. - checkpoint_workspace = ( - max(checkpoint_workspace, head_workspace_bytes, checkpoint_floor[1]) - + hybridep_growth_bytes + checkpoint_workspace = max( + checkpoint_workspace, head_workspace_bytes, checkpoint_floor[1] ) - required += int(hybridep_growth_bytes * _MEMORY_SAFETY_FACTOR) forward_required = required if gradient: required = max( @@ -4137,13 +4148,16 @@ def _subforward_cost( * _MEMORY_SAFETY_FACTOR ), ) + # HybridEP buffer growth stays allocated through the forward and + # backward peaks, but is not forward retention. return _SubforwardCost( - required=required, + required=required + int(hybridep_growth_bytes * _MEMORY_SAFETY_FACTOR), retained=retained, checkpoint_retained=checkpoint_retained, checkpoint_workspace=checkpoint_workspace, checkpoint_input_gradient=gradient, checkpoint_peak_increment=required - forward_required, + hybridep_growth=hybridep_growth_bytes, ) def _retained_memory_bytes( diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 56b17af34..aa30b54d8 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -596,8 +596,9 @@ def held(rows): base = rank._plan_cost(plan) charged = int(full * 1.1) assert grown.retained == base.retained < base.required - assert grown.checkpoint_workspace - base.checkpoint_workspace == full - assert grown.required - base.required in (charged, charged + 1) + assert grown.hybridep_growth == full and base.hybridep_growth == 0 + assert grown.checkpoint_workspace == base.checkpoint_workspace + assert grown.required - base.required == charged # A split pays it once, however many children would grow the buffer. split = rank._split_required_memory([grown, grown, grown]) extra = split - rank._split_required_memory([base, base, base]) @@ -605,3 +606,25 @@ def held(rows): rank.runtime.provider.expert_model_parallel_size = 1 monkeypatch.undo() assert rank._plan_hybridep_growth_bytes(plan) == 0 + + +def test_split_charges_the_largest_hybridep_growth_beside_any_child_peak(): + from art.trainer_rank._impl import _SubforwardCost + + def child(workspace: int, growth: int) -> _SubforwardCost: + peak = int((workspace + 1) * 1.1) + return _SubforwardCost( + required=peak + int(growth * 1.1), + retained=0, + checkpoint_workspace=workspace, + checkpoint_input_gradient=1, + hybridep_growth=growth, + ) + + # The first child grows the buffer most; the second has the larger + # workspace, which runs while that buffer is still allocated. + required = TrainerRank._split_required_memory( + [child(10_000, 10_000), child(15_000, 1_000)] + ) + assert required >= int((2 + 15_000 + 10_000) * 1.1) + assert required >= int((15_000 + 1) * 1.1) + int(10_000 * 1.1) From 587ae7e6b40568a41687a05388a41ac50cc1533c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 01:19:34 +0000 Subject: [PATCH 4/4] Add the HybridEP growth field to the replay fixture Replay compares recorded cost components with asdict(cost), which now includes hybridep_growth; a tampered value must also fail to match. Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_trainer_rank_planner_reports.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/unit/test_trainer_rank_planner_reports.py b/tests/unit/test_trainer_rank_planner_reports.py index 842755e5a..1adce5472 100644 --- a/tests/unit/test_trainer_rank_planner_reports.py +++ b/tests/unit/test_trainer_rank_planner_reports.py @@ -297,6 +297,7 @@ def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): "checkpoint_workspace": 0, "checkpoint_input_gradient": 0, "checkpoint_peak_increment": 0, + "hybridep_growth": 0, }, } ], @@ -377,6 +378,7 @@ def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): "checkpoint_input_gradient", "checkpoint_workspace", "checkpoint_peak_increment", + "hybridep_growth", ): altered = json.loads(path.read_bytes()) altered["replay"]["memory_replay"]["estimates"][0]["cost_components"][