Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
29 commits
Select commit Hold shift + click to select a range
f16b126
Raise EP routed-row pricing to each checkpoint's observed routing
bradhilton Sep 26, 2026
6cca90b
Agree on checkpoint targets; observe no-grad waves; seed snapshots
bradhilton Sep 26, 2026
e34327d
Also agree on the route epoch counter at loads; cover snapshot targets
bradhilton Sep 26, 2026
572be04
Read the route epoch counter with a default for bare trainers
bradhilton Sep 26, 2026
43db31f
Collect HybridEP managers into a typed list
bradhilton Sep 26, 2026
b4ce830
Merge the rebased per-rank CP2 layouts into routed-share observation
bradhilton Sep 26, 2026
4b87e8f
Merge the per-rank CP2 layout review fixes into routed-share observation
bradhilton Sep 26, 2026
3999370
Merge the layout-path TE test into routed-share observation
bradhilton Sep 27, 2026
b978418
Merge the layout-boundary adapter-gradient port (#978 53feda89b)
bradhilton Sep 28, 2026
5b6a947
Merge the #978 combine-floor test (5163520d8)
bradhilton Sep 28, 2026
66e945a
Merge the #978 per-rank adapter pairing (75a3912a7)
bradhilton Sep 28, 2026
59da6f8
Merge #978's sequential gradient groups (f47a06f27)
bradhilton Sep 28, 2026
1645975
Merge #978's pending-gradient tests (29c6a82ad)
bradhilton Sep 28, 2026
1768ab1
Merge #978's closed-form group order and pending reuse (500553bb4)
bradhilton Sep 28, 2026
e309eaa
Merge #978's two-slot split test (6eb86f7b4)
bradhilton Sep 28, 2026
dd4adfb
Merge #978's two-slot lower-bound assertion (bd26e7e9a)
bradhilton Sep 28, 2026
ca126c5
Merge #978's layer-mismatch test
bradhilton Sep 28, 2026
ea9ea84
Merge #978's staged head and decoder pricing
bradhilton Sep 28, 2026
1d79f15
Merge #978's head-stage row state
bradhilton Sep 28, 2026
44b622e
Merge #978's traced head staging
bradhilton Sep 28, 2026
c9b12b6
Merge #978's fused-statistics proof for head staging
bradhilton Sep 28, 2026
9dda1d6
Merge #978's TE-free profile readings
bradhilton Sep 28, 2026
0404e27
Merge #978's strict staged head statistics and V1t revert
bradhilton Sep 28, 2026
a8a3d5c
Merge #978's per-kernel fused-statistics proof
bradhilton Sep 28, 2026
e288179
Merge #978's strict binding per thread, group and staged admission
bradhilton Sep 28, 2026
b6bed2f
Raise EP routed-row pricing to each checkpoint's observed routing, on…
bradhilton Sep 29, 2026
b76895a
Merge #978's round-2 port (645c906e9) into the routed-share port
bradhilton Sep 29, 2026
726317d
Merge #978's ungrouped layouts test (bf79e4468)
bradhilton Sep 29, 2026
07bc450
Record #981's pre-extraction history (e28817973) as ported
bradhilton Sep 29, 2026
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
4 changes: 3 additions & 1 deletion .github/workflows/prek.yml
Original file line number Diff line number Diff line change
Expand Up @@ -248,6 +248,7 @@ jobs:
tests/unit/test_trainer_rank_converted_memory.py \
tests/unit/test_trainer_rank_layout_memory.py \
tests/unit/test_context_parallel_retained_bytes.py \
tests/unit/test_trainer_rank_routed_share.py \
tests/unit/test_trainer_rank_split.py \
tests/unit/test_megatron_compile_garbage.py \
tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_fits_after_recovery \
Expand Down Expand Up @@ -301,4 +302,5 @@ jobs:
--ignore=tests/unit/test_megatron_compile_garbage.py \
--ignore=tests/unit/test_trainer_rank_converted_memory.py \
--ignore=tests/unit/test_trainer_rank_layout_memory.py \
--ignore=tests/unit/test_context_parallel_retained_bytes.py
--ignore=tests/unit/test_context_parallel_retained_bytes.py \
--ignore=tests/unit/test_trainer_rank_routed_share.py
29 changes: 27 additions & 2 deletions src/art/trainer_rank/_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -1704,6 +1704,7 @@ def snapshot_checkpoint(trainer: TrainerRank, source: str, destination: str) ->
destination,
dict(source_slot.config),
source_slot.revision,
source_slot.route_epoch,
destination_slot is not None,
)
if any(value != identity for value in _gather(identity, group)):
Expand Down Expand Up @@ -1770,6 +1771,7 @@ def snapshot_checkpoint(trainer: TrainerRank, source: str, destination: str) ->
_restore_slots(model_snapshot)
trainer._checkpoint_slots.pop(destination, None)
raise
trainer._commit_route_epoch(destination, source=source_slot)
trainer._snapshot_checkpoint_names.add(destination)
return True

Expand All @@ -1782,7 +1784,13 @@ def discard_snapshot_checkpoint(trainer: TrainerRank, checkpoint: str) -> None:
trainer._default_slot_ref is not None
and trainer._default_slot_ref.name == checkpoint
) or any(ref.name == checkpoint for ref in trainer._slot_stack)
state = (slot is not None, False if slot is None else slot.snapshot, active)
state = (
checkpoint,
slot is not None,
False if slot is None else slot.snapshot,
None if slot is None else slot.route_epoch,
active,
)
if any(value != state for value in _gather(state, group)):
raise trainer._slot_state_error(
"Checkpoint snapshot state differs across ranks"
Expand Down Expand Up @@ -1815,6 +1823,7 @@ def discard_snapshot_checkpoint(trainer: TrainerRank, checkpoint: str) -> None:
_restore_slots(model_snapshot)
trainer._checkpoint_slots[checkpoint] = slot
raise
trainer._forget_route_epoch(slot)


def _commit_slot(trainer: TrainerRank, source: str, destination: str) -> None:
Expand Down Expand Up @@ -1973,10 +1982,25 @@ def load_checkpoint(
forward_only: bool = False,
) -> None:
group = _ensure_group(trainer)
if any(value != source.digest for value in _gather(source.digest, group)):
# Every rank must replace the same name, holding the same content, at the
# same point in its epoch sequence, so each gives the new content the same
# route epoch and forgets the same one.
current = trainer._checkpoint_slots.get(name)
target = (
source.digest,
name,
None if current is None else current.route_epoch,
getattr(trainer, "_route_epochs", 0),
)
targets = _gather(target, group)
if any(value[0] != source.digest for value in targets):
raise trainer._slot_state_error(
f"Checkpoint {name!r} content differs across ranks"
)
if any(value != target for value in targets):
raise trainer._slot_state_error(
f"Checkpoint {name!r} load target differs across ranks"
)
config = _phase(
lambda: trainer._validate_checkpoint_adapter_config(
name, source.config, alpha=None
Expand Down Expand Up @@ -2086,6 +2110,7 @@ def commit() -> None:
except BaseException:
_rollback_load(trainer, snapshot, temporary, name, previous, group)
raise
trainer._commit_route_epoch(name, previous)


def snapshot_prepared_checkpoint(
Expand Down
136 changes: 134 additions & 2 deletions src/art/trainer_rank/_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -892,6 +892,9 @@ class _CheckpointSlot:
custom: dict[str, _CustomObject] = dataclass_field(default_factory=dict)
custom_payload: "PreparedCustomPayload | None" = None
snapshot: bool = False
# Key for this committed content's routing observations; see
# ``TrainerRank._commit_route_epoch``.
route_epoch: int | None = None


@dataclass(frozen=True)
Expand Down Expand Up @@ -1598,6 +1601,13 @@ def _expert_lora_weight_storage(
# 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}
# Headroom above a checkpoint's observed routed share when that share exceeds
# the allowance. For the EP2 policy above, call-to-call change was 0.013 and
# batch resampling reached about 0.04 above the median.
_ROUTED_SHARE_MARGIN = 0.10
# Route epochs travel as float64 in the handoff exchange; stop observing well
# before two could round to the same value.
_ROUTE_EPOCH_LIMIT = 2**40


def _ep_routed_row_allowance(ep: int) -> float:
Expand Down Expand Up @@ -2170,6 +2180,11 @@ def memory_field(name: str, default: Any = None) -> Any:
self._slot_stack: list[LoRASlotRef] = []
self._checkpoint_slots: dict[str, _CheckpointSlot] = {}
self._snapshot_checkpoint_names: set[str] = set()
# Highest routed share observed per checkpoint route epoch, identical
# on every rank (``_release_cached_memory_for_backward``).
self._route_epochs = 0
self._routed_share_max: dict[int, float] = {}
self._last_routed_share: float | None = None
self._prepared_lora_exports: dict[str, tuple[str, _PreparedLoraExport]] = {}
self._checkpoint_prefetches: dict[str, Future[PreparedCheckpoint]] = {}
self._checkpoint_prefetch_sources: dict[str, str] = {}
Expand Down Expand Up @@ -2694,11 +2709,24 @@ def _release_cached_memory_for_backward(
# outputs. Forward has already executed: never replan or retry here.
with self._cache_recovery_episode(error=error) as (owner, started):
exchange_error: BaseException | None = None
# EP is the same on every rank, so the payload shape is too. Every
# rank must reach the exchange, whatever its observation does.
routing = getattr(getattr(self, "_parallel_shape", None), "ep", 1) > 1
share, epoch = -1.0, -1
if routing and error is None:
try:
share, epoch = self._local_routed_share(plan)
except Exception:
share, epoch = -1.0, -1
except BaseException as exc:
# Cancellation still exchanges, as a failed forward.
error = exc
try:
failed, gradients = self._recovery_reduce(
failed, gradients, *observed = self._recovery_reduce(
[
float(error is not None),
float(any(group.grad_enabled for group in plan.groups)),
*((share, float(epoch), -float(epoch)) if routing else ()),
],
op="MAX",
sync_across_dp=True,
Expand All @@ -2711,6 +2739,13 @@ def _release_cached_memory_for_backward(
raise self._memory_error_with_reduction_note(error, exchange_error)
if failed:
raise RuntimeError("Forward failed on another rank before handoff")
self._last_routed_share = None
if observed:
# Record only when every rank observed the same checkpoint
# epoch; any empty, unsupported or other-slot rank sends -1.
share, highest, lowest = observed
if highest >= 0 and highest == -lowest:
self._record_routed_share(int(highest), share)
if not gradients:
return
self._try_cache_recovery(
Expand All @@ -2721,6 +2756,55 @@ def _release_cached_memory_for_backward(
handoff_grad=any(group.grad_enabled for group in plan.groups),
)

def _local_routed_share(self, plan: _AnyForwardPlan) -> tuple[float, int]:
"""This rank's worst-layer routed share and checkpoint epoch, or -1s.

HybridEP keeps each MoE layer's received rows per local expert after
combine, so at the handoff of a flat plan with one group they are this
forward's dispatch, with or without gradients. The share is the most
loaded layer's received pairs over top-k times the balanced rows that
pricing scales (``_plan_group_balanced_rows``), so a share above the
allowance means more rows arrived than were priced.
"""
if (
not isinstance(plan, _FlatForwardPlan)
or len(plan.groups) != 1
or plan.groups[0].packed.tokens.numel() == 0
or plan.signature.topology[2] <= 1
or not getattr(self, "_ep_group_is_cp_group", False)
or not getattr(self, "_moe_memory_supported", False)
):
return -1.0, -1
epoch = self._route_epoch(plan.groups[0].slot_ref)
(balanced,) = self._plan_group_balanced_rows(plan)
if epoch is None or epoch >= _ROUTE_EPOCH_LIMIT or balanced <= 0:
return -1.0, -1
from megatron.core.transformer.moe.token_dispatcher import _HybridEPManager

managers: list[Any] = []
for chunk in self.runtime.model:
for module in chunk.modules():
dispatcher = getattr(module, "token_dispatcher", None)
manager = getattr(dispatcher, "_comm_manager", None)
if type(manager) is _HybridEPManager:
managers.append(manager)
if not managers or len(managers) != self._moe_layers:
return -1.0, -1
received = torch.stack(
[manager.tokens_per_expert.detach().sum() for manager in managers]
)
topk = managers[0].config.moe_router_topk
return int(received.max().item()) / (topk * balanced), epoch

def _record_routed_share(self, epoch: int, share: float) -> None:
shares = self._routed_share_max
shares[epoch] = max(share, shares.get(epoch, share))
self._last_routed_share = share
observation = getattr(self, "_planner_observation", None)
if observation is not None:
replay = observation["replay"]
observation["replay"] = lambda: {**replay(), "routed_share": share}

@overload
def dp_rank_forward(
self,
Expand Down Expand Up @@ -3183,7 +3267,9 @@ def last_forward_telemetry(self) -> dict[str, Any]:
are the admitted plan's memory check (for a split: every retained
graph plus the largest subforward's ephemeral share). A call refused
with ``TrainerRankMemoryError`` is still reflected, with the binding
check that refused it.
check that refused it. ``routed_share`` is the micro-batch's worst
observed expert-parallel routed share over its balanced rows, when
every rank observed the same checkpoint, and otherwise None.
"""

if self._last_forward_telemetry_snapshot is None:
Expand Down Expand Up @@ -4953,6 +5039,51 @@ def _cp_group_model_tokens(
),
)

def _route_epoch(self, ref: "LoRASlotRef | None") -> int | None:
if ref is None or ref.name is None:
return None
slot = getattr(self, "_checkpoint_slots", {}).get(ref.name)
return None if slot is None else slot.route_epoch

def _observed_routed_share(self, ref: "LoRASlotRef | None") -> float | None:
epoch = self._route_epoch(ref)
if epoch is None:
return None
return getattr(self, "_routed_share_max", {}).get(epoch)

def _commit_route_epoch(
self,
name: str,
previous: _CheckpointSlot | None = None,
*,
source: _CheckpointSlot | None = None,
) -> None:
"""Key a checkpoint's routing observations after every rank committed it.

Callers run this once the commit agreed across ranks, in the same order
everywhere, so each rank gives the same content the same new epoch.
Epochs are never reused. Loading over a name forgets the old content's
observations; optimizer steps keep the epoch, so its share only grows.
A snapshot starts from its ``source``'s share: same weights, same
routing.
"""
epoch = self._route_epochs
self._checkpoint_slots[name].route_epoch = epoch
self._route_epochs += 1
inherited = (
None
if source is None or source.route_epoch is None
else self._routed_share_max.get(source.route_epoch)
)
if inherited is not None:
self._routed_share_max[epoch] = inherited
if previous is not None:
self._forget_route_epoch(previous)

def _forget_route_epoch(self, slot: _CheckpointSlot) -> None:
if slot.route_epoch is not None:
self._routed_share_max.pop(slot.route_epoch, None)

_group_head_workspace_bytes = _memory._group_head_workspace_bytes
_layer_gdn_inputs = _memory._layer_gdn_inputs
_layout_checkpoint_floor = _memory._layout_checkpoint_floor
Expand Down Expand Up @@ -4981,6 +5112,7 @@ def _cp_group_model_tokens(
_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_group_balanced_rows = _micro_batch_planner._plan_group_balanced_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
Expand Down
34 changes: 34 additions & 0 deletions src/art/trainer_rank/_micro_batch_planner.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
replace,
)
import hashlib
import math
from typing import TYPE_CHECKING

from art.trainer_rank import _impl
Expand Down Expand Up @@ -610,6 +611,28 @@ def _plan_head_backward_traced(self: TrainerRank, plan: _FlatForwardPlan) -> boo

def _plan_group_routed_rows(
self: TrainerRank, plan: _FlatForwardPlan
) -> tuple[int, ...]:
"""Balanced rows per group, raised to each slot's observed routing.

The cold allowance prices top-k x allowance pairs per balanced row. A
checkpoint whose worst observed share plus a margin exceeds that is
priced at the higher share instead; nothing is ever priced lower.
"""
balanced = self._plan_group_balanced_rows(plan)
shares = [self._observed_routed_share(group.slot_ref) for group in plan.groups]
if all(share is None for share in shares):
return balanced
cold = _impl._ep_routed_row_allowance(self._parallel_shape.ep)
return tuple(
rows
if share is None or share + _impl._ROUTED_SHARE_MARGIN <= cold
else math.ceil(rows * (share + _impl._ROUTED_SHARE_MARGIN) / cold)
for rows, share in zip(balanced, shares, strict=True)
)


def _plan_group_balanced_rows(
self: TrainerRank, plan: _FlatForwardPlan
) -> tuple[int, ...]:
"""Rows one rank's experts receive per group at balanced routing.

Expand Down Expand Up @@ -771,6 +794,7 @@ def _split_request_order(

def _reset_planning_telemetry(self: TrainerRank) -> None:
self._planning_seconds_accum = 0.0
self._last_routed_share = None
with self._layout_cache_lock:
self._speculative_planning_seconds = 0.0

Expand All @@ -785,6 +809,9 @@ def _snapshot_planning_telemetry(
if isinstance(plan, _impl._SplitForwardPlan)
else (tuple(range(plan.request_count)),)
)
# Consume the share so a later forward without a handoff reports none.
routed_share = getattr(self, "_last_routed_share", None)
self._last_routed_share = None
self._last_forward_telemetry_snapshot = {
"planning_ms": self._planning_seconds_accum * 1_000.0,
"speculative_planning_ms": speculative_seconds * 1_000.0,
Expand All @@ -795,6 +822,7 @@ def _snapshot_planning_telemetry(
"subforward_request_indices": partition,
"predicted_peak_bytes": check.estimated_required_bytes,
"usable_limit_bytes": check.available_bytes,
"routed_share": routed_share,
}


Expand Down Expand Up @@ -1631,6 +1659,12 @@ def _fill_planner_snapshot(
self._plan_hybridep_growth_bytes(child)
),
},
# Each group's observed share behind its routed rows;
# None prices the cold allowance.
"routed_share_max": [
self._observed_routed_share(group.slot_ref)
for group in child.groups
],
"expected_required_bytes": cost.required,
"retained_bytes": cost.retained,
"cost_components": asdict(cost),
Expand Down
3 changes: 3 additions & 0 deletions src/art/trainer_rank/_planner_replay.py
Original file line number Diff line number Diff line change
Expand Up @@ -149,6 +149,9 @@ def capture(rank: Any, plan: Any) -> dict[str, Any]:
"_recomputed_mixer_widths",
# Producers of recorded arguments replay checks against the facts.
"_plan_group_routed_rows",
"_plan_group_balanced_rows",
"_observed_routed_share",
"_route_epoch",
"_plan_head_backward_traced",
"_head_backward_traced",
):
Expand Down
Loading
Loading