diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 5f0006c4a..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: @@ -1494,6 +1497,48 @@ 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_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. + + ``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 = _hybridep_rows_per_rank(capacity, ranks) * ranks + return tokens * (2 * hidden + 5 * experts + 4 * (hidden // 128)) + + def _moe_output_bytes_per_token( model: Sequence[torch.nn.Module], shape: ParallelShape, @@ -1503,8 +1548,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 +1561,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 = _EP_ROUTED_ROW_ALLOWANCE if shape.ep > 1 else 1 coefficient = 0 for chunk in model: for layer in chunk.modules(): @@ -1536,15 +1588,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 +1664,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 +1689,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) @@ -3287,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 @@ -3794,6 +3862,49 @@ 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) + ) + etp = int(getattr(provider, "expert_tensor_parallel_size", 1) or 1) + ranks = ep * etp + if _hybridep_rows_per_rank(capacity, ranks) <= held: + return 0 + # 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], ...]: """Physical rows per group on the most loaded context-parallel rank.""" topology = self._topology() if plan.signature.topology[2] > 1 else None @@ -3972,6 +4083,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 +4099,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, @@ -4035,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( @@ -6097,6 +6213,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,7 +6941,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), - ) + ) + 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: diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 33b8ce9df..aa30b54d8 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -188,6 +188,54 @@ 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 + + +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)] +) +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 +540,91 @@ 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): + 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 + 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(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(512), full), (held(1024), 0)): + monkeypatch.setattr(fused_a2a, "_hybrid_ep_buffer", current) + assert rank._plan_hybridep_growth_bytes(plan) == growth + # 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) + charged = int(full * 1.1) + assert grown.retained == base.retained < base.required + 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]) + assert charged - 1 <= extra <= charged + 1 + 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) 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"][