Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
19 commits
Select commit Hold shift + click to select a range
f0670b9
Price the adapter gradients a recompute backward holds at its peak
bradhilton Sep 27, 2026
ecf2ce2
Count head-tied and custom checkpoint gradients as live throughout
bradhilton Sep 27, 2026
067048e
Order slot names safely and leave base-model groups out of adapter slots
bradhilton Sep 27, 2026
c35b238
Test that base-model gradient groups price and defer nothing
bradhilton Sep 28, 2026
f1fe9e7
Merge main (762c89dfa) into the adapter-gradient floor
bradhilton Sep 28, 2026
e54f4bc
Walk gradient groups' backward one after another, in the worst order
bradhilton Sep 28, 2026
0df629c
Price gradient groups' worst backward order in closed form
bradhilton Sep 28, 2026
2a64195
Test that a layer-count mismatch prices no adapter-gradient extra
bradhilton Sep 28, 2026
f0e11e0
Merge remote-tracking branch 'origin/main' into dalinar/adapter-grad-…
bradhilton Sep 28, 2026
d3093bf
Merge remote-tracking branch 'origin/main' into dalinar/adapter-grad-…
bradhilton Sep 29, 2026
be1f15d
Replay the adapter-gradient floor from frozen slot facts
bradhilton Sep 29, 2026
a9135c1
Merge remote-tracking branch 'origin/main' into dalinar/adapter-grad-…
bradhilton Sep 29, 2026
b8a2fd1
Narrow the replay capture's slot name before sizing it
bradhilton Sep 29, 2026
6b1b918
Bound replay's recorded layer count and tighten adapter facts
bradhilton Sep 29, 2026
eedef3d
Patch the reader with monkeypatch in the custom-reader test
bradhilton Sep 29, 2026
af503a0
Merge remote-tracking branch 'origin/main' into dalinar/adapter-grad-…
bradhilton Sep 29, 2026
a4f05f8
Validate adapter pending facts before scanning them
bradhilton Sep 29, 2026
3fa5dbf
Refuse mixed slot kinds in facts and a zero layer count at capture
bradhilton Sep 29, 2026
c41148e
Name the refusal each forged adapter fact must hit
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
2 changes: 2 additions & 0 deletions .github/workflows/prek.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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 \
Expand Down Expand Up @@ -282,6 +283,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 \
Expand Down
50 changes: 49 additions & 1 deletion src/art/trainer_rank/_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
from dataclasses import dataclass, replace
from dataclasses import field as dataclass_field
from functools import partial
import json
import logging
import math
import os
Expand Down Expand Up @@ -124,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

Expand Down Expand Up @@ -1060,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`` 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: str = ""

@property
def ephemeral(self) -> int:
Expand Down Expand Up @@ -2940,18 +2952,33 @@ 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 = self._gradient_slots(group_rows, slot_refs)
adapter_gradient = (
self._checkpoint_adapter_gradient_bytes(
self._checkpoint_gradient_groups(group_rows, slot_refs)
)
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
),
)
Expand All @@ -2965,6 +2992,22 @@ 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=json.dumps(
[
[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 "",
)

def last_forward_telemetry(self) -> dict[str, Any]:
Expand Down Expand Up @@ -4704,6 +4747,11 @@ def _gather_tensor_parallel_logits(self, logits: torch.Tensor) -> torch.Tensor:
_checkpoint_moe_bytes_per_token = _memory._checkpoint_moe_bytes_per_token
_moe_workspace_bytes = _memory._moe_workspace_bytes
_checkpoint_memory_floor = _memory._checkpoint_memory_floor
_gradient_slots = staticmethod(_memory._gradient_slots)
_pending_adapter_gradient_bytes = _memory._pending_adapter_gradient_bytes
_checkpoint_gradient_groups = _memory._checkpoint_gradient_groups
_checkpoint_adapter_gradient_bytes = _memory._checkpoint_adapter_gradient_bytes
_adapter_gradient_walk = staticmethod(_memory._adapter_gradient_walk)
_retained_memory_bytes = _memory._retained_memory_bytes
_estimate_flat_forward = _memory._estimate_flat_forward
_update_peak_memory_profile = _memory._update_peak_memory_profile
Expand Down
221 changes: 219 additions & 2 deletions src/art/trainer_rank/_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@

from __future__ import annotations

from collections.abc import Iterable, Sequence
from collections.abc import Iterable, Iterator, Sequence
from contextlib import nullcontext
import hashlib
import math
Expand Down Expand Up @@ -45,12 +45,32 @@ def _split_required_memory(costs: Sequence[_impl._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 * _impl._MEMORY_SAFETY_FACTOR))
Expand Down Expand Up @@ -625,6 +645,182 @@ def _checkpoint_floor_from_facts(
return retained, workspace


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: TrainerRank, 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 = _impl._language_model(chunk).decoder.layers
except (AttributeError, RuntimeError):
return ()
layer_of: dict[int, int] = {}
params: dict[int, torch.nn.Parameter] = {}

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(module, LoRA) and "_slot_keys" in vars(module):
for ref in refs:
yield from module.lora_slot_params(ref)

for index, layer in enumerate(layers):
for param in slot_params(layer.modules()):
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 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)
# 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 (
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_gradient_groups(
self: TrainerRank,
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: TrainerRank, groups: Sequence[tuple[LoRASlotRef | None, Sequence[int]]]
) -> 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``).
"""
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))
return self._adapter_gradient_walk(chains)


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 (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 _retained_memory_bytes(
self: TrainerRank,
signature: _impl._MemorySignature,
Expand Down Expand Up @@ -705,6 +901,19 @@ 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 and ref.name 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 _impl._gdn_memory.model_shapes(self) is not None
Expand Down Expand Up @@ -1169,11 +1378,19 @@ def _estimate_required_memory_bytes_from_values(
if checkpoint_memory is None
else checkpoint_memory
)
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(
self._checkpoint_gradient_groups(group_rows, slot_refs)
)
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])
+ (retained if include_checkpoint_input_gradient else 0),
+ backward,
)
if signature.topology[2] > 1:
# Local head results coexist with full CP outputs during gathering.
Expand Down
3 changes: 3 additions & 0 deletions src/art/trainer_rank/_planner_misses.py
Original file line number Diff line number Diff line change
Expand Up @@ -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])
Expand Down
Loading
Loading