Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
167 changes: 143 additions & 24 deletions src/art/trainer_rank/_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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():
Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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)
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand All @@ -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,
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down
Loading
Loading