diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 924abc0cc..f7360485f 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -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 \ @@ -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 \ diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index a25505c53..8fcf2497a 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -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 @@ -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 @@ -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: @@ -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 ), ) @@ -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]: @@ -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 diff --git a/src/art/trainer_rank/_memory.py b/src/art/trainer_rank/_memory.py index 70e3c9250..40ea5b5b7 100644 --- a/src/art/trainer_rank/_memory.py +++ b/src/art/trainer_rank/_memory.py @@ -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 @@ -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)) @@ -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, @@ -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 @@ -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. diff --git a/src/art/trainer_rank/_planner_misses.py b/src/art/trainer_rank/_planner_misses.py index bb7e8297c..b1df930f1 100644 --- a/src/art/trainer_rank/_planner_misses.py +++ b/src/art/trainer_rank/_planner_misses.py @@ -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]) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index f0da8bbe1..c3b7db458 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -11,16 +11,24 @@ from dataclasses import asdict import json from types import MethodType, SimpleNamespace -from typing import Any +from typing import Any, NamedTuple from . import _gdn_memory, _impl, _memory _MAX_BYTES = 262144 _MAX_GROUPS = 1024 +# Replay sizes per-layer tuples from this recorded rank field; bound it as capture does. +MAX_LAYERS = 1024 _MAX_SEGMENTS = 4096 _MAX_INPUT_VALUES = 1_000_000 +# Estimators bound on TrainerRank as plain functions, not methods. +_STATIC_ESTIMATORS = frozenset( + {"_split_required_memory", "_gradient_slots", "_adapter_gradient_walk"} +) + + _REFUSALS = frozenset( { "runtime_group_inventory_over_limit", @@ -68,7 +76,10 @@ def reserve(size: int) -> None: def capture(rank: Any, plan: Any) -> dict[str, Any]: - if rank._num_layers > 1024 or not 0 < len(plan.groups) <= _MAX_GROUPS: + if ( + not 0 < rank._num_layers <= MAX_LAYERS + or not 0 < len(plan.groups) <= _MAX_GROUPS + ): raise ValueError("runtime_group_inventory_over_limit") if ( getattr(rank.runtime.provider, "expert_model_parallel_size", 1) > 1 @@ -96,12 +107,17 @@ def capture(rank: Any, plan: Any) -> dict[str, Any]: "_physical_tokens", "_plan_group_rows", "_plan_retained_tokens", + "_gradient_slots", + "_pending_adapter_gradient_bytes", + "_checkpoint_gradient_groups", + "_checkpoint_adapter_gradient_bytes", + "_adapter_gradient_walk", ): method = getattr(rank, name) expected = getattr(_impl.TrainerRank, name) supported = ( method is expected - if name == "_split_required_memory" # The sole static estimator. + if name in _STATIC_ESTIMATORS else type(method) is MethodType and method.__self__ is rank and method.__func__ is expected @@ -214,6 +230,18 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: if backwards else 0 ) + adapter = None + if group.grad_enabled and name is not None: + # The live estimator reads this slot's unallocated gradient bytes + # per decoder layer; freeze them with the selection. + kind = getattr(group.slot_ref, "kind", None) + if kind is not None and (type(kind) is not str or len(kind) > 64): + raise ValueError("runtime_slot_identity_unsupported") + pending = rank._pending_adapter_gradient_bytes((group.slot_ref,)) + if len(pending) > 1025: + raise ValueError("runtime_shape_inventory_over_limit") + reserve(128 + 12 * len(name) + 24 * len(pending)) + adapter = {"kind": kind, "name": name, "pending": [int(v) for v in pending]} model = _gdn_memory.model_shapes(rank, group.slot_ref) if has_grad else None if model is not None: if len(model[1]) > 1024: @@ -231,6 +259,7 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: "gradient": terms(True, group.slot_ref), "head_rows": projected, "head_target_rows": target_rows, + "adapter": adapter, "gdn": None if model is None else { @@ -253,7 +282,7 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: } ) facts = { - "version": 1, + "version": 2, "checkpoint_layers": _memory._checkpoint_layers( rank, rank._plan_group_rows(plan) ), @@ -290,7 +319,7 @@ def integer(value: Any, *, minimum: int = 0) -> None: "groups", }, ) - if type(facts["version"]) is not int or facts["version"] != 1: + if type(facts["version"]) is not int or facts["version"] != 2: raise ValueError("unsupported runtime facts version") for key in ( "checkpoint_layers", @@ -319,6 +348,7 @@ def integer(value: Any, *, minimum: int = 0) -> None: "gradient", "head_rows", "head_target_rows", + "adapter", "gdn", }, ) @@ -363,6 +393,25 @@ def integer(value: Any, *, minimum: int = 0) -> None: raise ValueError("invalid MoE stage") for value in stage: integer(value) + adapter = group["adapter"] + if adapter is not None: + fields(adapter, {"kind", "name", "pending"}) + if ( + not group["grad"] + or (adapter["kind"] is not None and type(adapter["kind"]) is not str) + or len(adapter["kind"] or "") > 64 + or type(adapter["name"]) is not str + or len(adapter["name"]) > 4096 + or type(adapter["pending"]) is not list + or len(adapter["pending"]) > 1025 + ): + raise ValueError("invalid adapter gradient facts") + reserve(128 + 12 * len(adapter["name"]) + 24 * len(adapter["pending"])) + for value in adapter["pending"]: + integer(value) + # A slot without a kind (megatron-less reference) has none pending. + if adapter["kind"] is None and any(adapter["pending"]): + raise ValueError("invalid adapter gradient facts") gdn = group["gdn"] if gdn is not None: fields(gdn, {"layers", "shapes", "segments"}) @@ -390,6 +439,14 @@ def integer(value: Any, *, minimum: int = 0) -> None: ) for value in segment.values(): integer(value) + # Live slot references all have a kind, or (without megatron) none do. + kinds = { + group["adapter"]["kind"] is None + for group in groups + if group["adapter"] is not None + } + if len(kinds) > 1: + raise ValueError("invalid adapter gradient facts") if len(json.dumps(facts, separators=(",", ":"))) > _MAX_BYTES: raise ValueError("runtime_facts_over_limit") @@ -420,11 +477,54 @@ def count(value: Any, depth: int = 0) -> None: count(value) +class _ReplaySlot(NamedTuple): + """A frozen adapter slot identity (the live LoRASlotRef's kind and name).""" + + kind: str | None + name: str + + class ReplayRank(_impl.TrainerRank): """The real estimator with runtime metadata readers replaced by frozen facts.""" _facts: dict[str, Any] | None = None + def _replay_slots(self, slot_refs: Any) -> Any: + # Replay passes each group's index; map it to that group's frozen slot. + if self._facts is None or slot_refs is None: + return slot_refs + groups = self._facts["groups"] + return tuple( + None + if (adapter := groups[index]["adapter"]) is None + else _ReplaySlot(adapter["kind"], adapter["name"]) + for index in slot_refs + ) + + def _gradient_slots(self, group_rows: Any, slot_refs: Any) -> Any: + return _memory._gradient_slots(group_rows, self._replay_slots(slot_refs)) + + def _checkpoint_gradient_groups(self, group_rows: Any, slot_refs: Any) -> Any: + return _memory._checkpoint_gradient_groups( + self, group_rows, self._replay_slots(slot_refs) + ) + + def _pending_adapter_gradient_bytes(self, refs: Any) -> tuple[int, ...]: + if self._facts is None: + return _memory._pending_adapter_gradient_bytes(self, refs) + refs = tuple(dict.fromkeys(refs)) + if not refs: + return () + if len(refs) != 1: + raise ValueError("replayed adapter gradients are frozen per slot") + for group in self._facts["groups"]: + adapter = group["adapter"] + if adapter is not None and (adapter["kind"], adapter["name"]) == tuple( + refs[0] + ): + return tuple(adapter["pending"]) + return () + def _head_workspace_bytes(self, rows: int) -> int: assert self._facts is not None return _memory._dense_head_bytes(self._facts["head_vocabulary"], rows) diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 575fafa54..c0c4b1ba8 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -171,6 +171,127 @@ def test_selected_slot_terms_are_replayed_and_frozen(layer, tmp_path): ) +def adapter_report(layer, tmp_path): + """A grouped report whose gradient slot has pending adapter gradients.""" + from test_trainer_rank_adapter_gradient_memory import lora, parameter + from test_trainer_rank_converted_memory import weights + from test_trainer_rank_pending_memory import rank_with_moe + from test_trainer_rank_slot_memory import load_slot + from test_trainer_rank_slot_memory import request as slot_request + + rank, _ = rank_with_moe(weights(layer, 8)) + load_slot(rank, "small", 1) + load_slot(rank, "large", 64) + # Unallocated slot gradients the recompute backward will allocate: small, + # distinct per-layer sizes (the fixture has 40 layers; keep this bounded). + layers = tr._language_model(rank.runtime.model[0]).decoder.layers + assert len(layers) <= 64 + sizes = [16 * (index + 1) for index in range(len(layers))] + assert sum(sizes) * 2 <= 66_560 # BF16 bytes, checked before allocating. + params = [] + for size, block in zip(sizes, layers, strict=True): + params.append(parameter(size)) + block.add_module("adapter", lora(large=[params[-1]])) + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + plan = rank._plan_flat_forward( + [slot_request("small", rows=2), slot_request("large", rows=65, grad=True)], + ensure_slots=False, + ) + original, costs = emitted(rank, plan, tmp_path) + return original, costs, params + + +def test_selected_adapter_gradients_are_replayed_and_frozen(layer, tmp_path): + original, costs, params = adapter_report(layer, tmp_path) + assert costs[0].checkpoint_adapter_gradient > 0 + groups = original["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "groups" + ] + assert groups[0]["adapter"] is None + assert groups[1]["adapter"]["name"] == "large" and any( + groups[1]["adapter"]["pending"] + ) + actual = reports.replay(original) + assert actual["aggregate"]["matches"] + assert all(item["matches"] for item in actual["estimates"]) + # Gradients allocated after selection cannot change the replayed answer. + for param in params: + param.grad = torch.zeros_like(param) + assert reports.replay(original) == actual + changed = deepcopy(original) + changed["replay"]["memory_replay"]["estimates"][0]["runtime_facts"]["groups"][1][ + "adapter" + ]["pending"][0] += 10**12 + result = reports.replay(changed) + assert not result["estimates"][0]["matches"] + assert ( + result["estimates"][0]["required_bytes"] + > actual["estimates"][0]["required_bytes"] + ) + + +@pytest.mark.parametrize( + "change", + ["gradient", "length", "value", "kind_length", "kindless_pending", "mixed_kinds"], +) +def test_adapter_fact_validation_rejects_forged_input(change, layer, tmp_path): + report, _, _ = adapter_report(layer, tmp_path) + group = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "groups" + ][1] + adapter = group["adapter"] + if change == "gradient": + group["grad"] = False + elif change == "length": + adapter["pending"] = [0] * 1026 + elif change == "value": + adapter["pending"][0] = 1.5 + elif change == "kind_length": + adapter["kind"] = "k" * 65 + elif change == "kindless_pending": + adapter["kind"] = None + else: + groups = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "groups" + ] + groups[0]["grad"] = True + groups[0]["adapter"] = {"kind": None, "name": "base", "pending": []} + # Name the refusal: a forged fact must fail validation, not a later check. + message = ( + "invalid runtime dimension" + if change == "value" + else "invalid adapter gradient facts" + ) + with pytest.raises(ValueError, match=message): + reports.replay(report) + + +@pytest.mark.parametrize("layers", [2**10 + 1, 0, 40.0]) +def test_replay_bounds_the_recorded_layer_count(layers, pending_rank, tmp_path): + rank = pending_rank + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, _ = emitted( + rank, rank._plan_flat_forward([request(65, grad=True)]), tmp_path + ) + # Replay sizes per-layer tuples from this field; refuse it before costing. + report["replay"]["memory_replay"]["rank"]["num_layers"] = layers + with pytest.raises(ValueError, match="layer count"): + reports.replay(report) + + +def test_custom_adapter_gradient_reader_is_explicitly_incomplete(monkeypatch): + from art.trainer_rank import _planner_replay + + rank = _rank(monkeypatch) + plan = rank._plan_flat_forward([_request(1)]) + original = rank._pending_adapter_gradient_bytes + monkeypatch.setattr( + rank, "_pending_adapter_gradient_bytes", lambda refs: original(refs) + ) + with pytest.raises(ValueError, match="custom_runtime_estimator"): + _planner_replay.capture(rank, plan) + + @pytest.mark.parametrize( "change", ["version", "group", "layout", "gdn_segment", "budget"] ) @@ -184,7 +305,7 @@ def test_fact_validation_rejects_inconsistent_or_unbounded_input( ) facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] if change == "version": - facts["version"] = 2 + facts["version"] += 1 elif change == "group": facts["groups"][0]["rows"] += 1 elif change == "layout": diff --git a/tests/unit/test_planner_replay_owner_budget.py b/tests/unit/test_planner_replay_owner_budget.py index 5292cf743..248683dd2 100644 --- a/tests/unit/test_planner_replay_owner_budget.py +++ b/tests/unit/test_planner_replay_owner_budget.py @@ -47,7 +47,7 @@ def test_shared_inventory_preflight_precedes_layout_construction( elif case == "invalid_facts": report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ "version" - ] = 2 + ] += 1 reason = "unsupported runtime facts version" else: monkeypatch.setattr(runtime, "_MAX_INPUT_VALUES", 10) diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py new file mode 100644 index 000000000..0da30ce0b --- /dev/null +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -0,0 +1,455 @@ +"""Adapter gradients at the recompute backward's peak: CPU admission math. + +Qwen3.6-35B-A3B CP2 allocator traces: a short first wave peaks at layer 0 with +830-900 MB of expert LoRA gradients live, a long one at the last layer. +""" + +from collections.abc import Sequence +import itertools +import random + +import pytest +from test_trainer_rank_checkpoint_memory import rank, requests +import torch + +from art.megatron.lora import LoRA, LoRASlotRef +from art.trainer_rank import TrainerRank +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) +from art.trainer_rank._impl import _MemoryProfile, _SubforwardCost + +POLICY = LoRASlotRef("checkpoint", "policy") +OTHER = LoRASlotRef("checkpoint", "other") + + +def oracle(pending, boundaries): + """Live bytes beyond the floor while backward recomputes each layer.""" + layers = len(boundaries) + return max( + 0, + *( + sum(pending[index:layers]) + pending[layers] - sum(boundaries[index + 1 :]) + for index in range(layers) + ), + ) + + +def with_pending(monkeypatch, r, pending): + # Like the real method: no slots, nothing pending. + monkeypatch.setattr( + r, + "_pending_adapter_gradient_bytes", + lambda refs: tuple(pending) if tuple(refs) else (), + ) + + +@pytest.mark.parametrize( + "pending, boundaries", + [ + # Uniform gradients above the boundaries: layer 0 is the peak. + ([23] * 4 + [0], [4] * 4), + # Uniform gradients below the boundaries: the last layer is the peak. + ([3] * 4 + [0], [10] * 4), + # Attention and GDN layers differ; the peak is interior to neither end. + ([0, 100, 0, 100, 5], [60, 60, 60, 60]), + # Uneven boundaries, gradients outside the decoder live throughout. + ([7, 0, 50, 1, 9], [1, 80, 2, 3]), + ], +) +def test_extra_is_the_worst_backward_layer_over_real_sizes( + monkeypatch, pending, boundaries +): + r = rank() + with_pending(monkeypatch, r, pending) + assert r._checkpoint_adapter_gradient_bytes(((POLICY, boundaries),)) == oracle( + pending, boundaries + ) + + +def test_extra_matches_every_layer_for_random_sizes(monkeypatch): + r = rank() + generator = random.Random(0) + for _ in range(200): + layers = generator.randint(1, 12) + pending = [ + generator.choice((0, generator.randint(0, 50))) for _ in range(layers + 1) + ] + boundaries = [generator.randint(0, 40) for _ in range(layers)] + with_pending(monkeypatch, r, pending) + extra = r._checkpoint_adapter_gradient_bytes(((POLICY, boundaries),)) + assert extra == oracle(pending, boundaries) + + +def test_no_pending_gradients_add_nothing(monkeypatch): + r = rank() + with_pending(monkeypatch, r, ()) + assert r._checkpoint_adapter_gradient_bytes(((POLICY, [4] * 40),)) == 0 + + +def test_a_layer_count_mismatch_prices_no_extra(monkeypatch): + r = rank() + with_slot_pending( + monkeypatch, r, {(POLICY,): [5] * 4 + [0], (OTHER,): [9] * 3 + [0]} + ) + # Pending gradients and boundaries must describe the same decoder layers; + # otherwise the whole term is left out rather than misplaced. + assert r._checkpoint_adapter_gradient_bytes(((POLICY, [1] * 4),)) > 0 + assert r._checkpoint_adapter_gradient_bytes(((OTHER, [1] * 4),)) == 0 + assert ( + r._checkpoint_adapter_gradient_bytes(((POLICY, [1] * 4), (OTHER, [1] * 4))) == 0 + ) + + +def sequential_oracle(chains): + """Worst live bytes beyond the floor, one group's backward after another. + + While a group recomputes layer i, groups run before it hold all their + gradients and none of their boundaries; groups yet to run hold all their + boundaries; the running group releases its boundaries above i. + """ + worst = 0 + for order in itertools.permutations(chains): + for position, (pending, boundaries) in enumerate(order): + done = order[:position] + layers = len(boundaries) + for index in range(layers): + gradients = ( + sum(sum(p) for p, _ in done) + + pending[layers] + + sum(pending[index:layers]) + ) + released = sum(sum(b) for _, b in done) + sum(boundaries[index + 1 :]) + worst = max(worst, gradients - released) + return worst + + +def with_slot_pending(monkeypatch, r, by_slot): + monkeypatch.setattr( + r, + "_pending_adapter_gradient_bytes", + lambda refs: by_slot.get(tuple(refs), ()), + ) + + +def test_gradient_groups_run_their_backward_one_after_another(monkeypatch): + r = rank() + # A short policy group beside a long group of another slot: whichever + # runs first, the other keeps its boundaries (or its gradients) meanwhile. + policy = ([30] * 4 + [0], [2] * 4) + other = ([5] * 4 + [1], [40] * 4) + with_slot_pending(monkeypatch, r, {(POLICY,): policy[0], (OTHER,): other[0]}) + extra = r._checkpoint_adapter_gradient_bytes( + ((POLICY, policy[1]), (OTHER, other[1])) + ) + assert extra == sequential_oracle([policy, other]) + # One chain with every group's boundaries released together would have + # priced far less. + combined = [a + b for a, b in zip(policy[0], other[0])] + assert extra > oracle(combined, [a + b for a, b in zip(policy[1], other[1])]) + # A base-model group owns no gradients, but its boundaries stay live + # while the policy group runs first. + base = ([0] * 5, [40] * 4) + assert ( + r._checkpoint_adapter_gradient_bytes(((POLICY, policy[1]), (None, base[1]))) + == sequential_oracle([policy, base]) + == oracle(*policy) + ) + + +def test_sequential_groups_match_every_order_for_random_sizes(monkeypatch): + r = rank() + generator = random.Random(1) + slots = [LoRASlotRef("checkpoint", f"slot{index}") for index in range(5)] + for _ in range(200): + layers = generator.randint(1, 6) + chains, groups, by_slot = [], [], {} + for slot in slots[: generator.randint(1, 5)]: + pending = [generator.randint(0, 30) for _ in range(layers + 1)] + boundaries = [generator.randint(0, 30) for _ in range(layers)] + if generator.random() < 0.25: + slot, pending = None, [0] * (layers + 1) + else: + by_slot[(slot,)] = pending + chains.append((pending, boundaries)) + groups.append((slot, boundaries)) + with_slot_pending(monkeypatch, r, by_slot) + assert r._checkpoint_adapter_gradient_bytes(groups) == sequential_oracle(chains) + + +def test_many_gradient_groups_price_exactly(monkeypatch): + r = rank() + slots = [LoRASlotRef("checkpoint", f"slot{index}") for index in range(6)] + # Only groups whose gradients outweigh their boundaries raise another + # group's peak by having run first. + chains = [([3, 1, 2], [100, 100])] * 3 + [([90, 40, 5], [2, 1])] * 3 + with_slot_pending( + monkeypatch, r, {(slot,): pending for slot, (pending, _) in zip(slots, chains)} + ) + groups = [(slot, boundaries) for slot, (_, boundaries) in zip(slots, chains)] + assert r._checkpoint_adapter_gradient_bytes(groups) == sequential_oracle(chains) + assert sequential_oracle(chains) < sum(sum(pending) for pending, _ in chains) + + +def lora(**slots: list[torch.nn.Parameter]) -> LoRA: + module = LoRA.__new__(LoRA) + torch.nn.Module.__init__(module) + module._slot_keys = {} + by_name = { + LoRASlotRef("checkpoint", name): params for name, params in slots.items() + } + module.lora_slot_params = lambda ref: by_name.get(ref, []) # type: ignore[method-assign] + return module + + +def parameter(elements: int, *, dtype=torch.bfloat16) -> torch.nn.Parameter: + return torch.nn.Parameter(torch.zeros(elements, dtype=dtype)) + + +def adapter_rank( + layers: Sequence[torch.nn.Module], outside: torch.nn.Module +) -> TrainerRank: + r = rank() + block = r.runtime.model[0].decoder + block.layers = torch.nn.ModuleList(layers) + r.runtime.model[0].head = outside + return r + + +def test_pending_gradients_count_local_unallocated_slot_parameters(): + shared = parameter(10) + allocated = parameter(1000) + allocated.grad = torch.zeros_like(allocated) + master = parameter(1000) + setattr(master, "main_grad", torch.zeros(1000)) + frozen = parameter(1000) + frozen.requires_grad_(False) + layers = [ + lora(policy=[parameter(3), shared]), + lora(policy=[parameter(5)], other=[parameter(1000)]), + lora(policy=[allocated, master, frozen]), + lora(policy=[shared, parameter(4, dtype=torch.float32)]), + ] + r = adapter_rank(layers, lora(policy=[parameter(6)])) + pending = r._pending_adapter_gradient_bytes([POLICY]) + # BF16 bytes per layer; the shared parameter counts once, at its highest + # layer; allocated, main-grad and frozen parameters are not pending; the + # parameter outside the decoder is live throughout. + assert pending == (3 * 2, 5 * 2, 0, (10 + 0) * 2 + 4 * 4, 6 * 2) + assert r._pending_adapter_gradient_bytes([POLICY, OTHER]) == ( + 3 * 2, + 5 * 2 + 1000 * 2, + 0, + 10 * 2 + 4 * 4, + 6 * 2, + ) + assert r._pending_adapter_gradient_bytes([LoRASlotRef("checkpoint", "x")]) == () + assert r._pending_adapter_gradient_bytes([]) == () + + +def test_a_parameter_the_head_also_uses_is_live_throughout(): + tied = parameter(7) + layers = [lora(policy=[tied, parameter(3)]), lora(policy=[parameter(5)])] + r = adapter_rank(layers, lora(policy=[tied])) + # The head's backward runs before the decoder's and allocates it first. + assert r._pending_adapter_gradient_bytes([POLICY]) == (3 * 2, 5 * 2, 7 * 2) + + +def test_a_module_the_head_also_uses_is_live_throughout(): + shared = lora(policy=[parameter(7)]) + r = adapter_rank([shared, torch.nn.Module()], shared) + assert r._pending_adapter_gradient_bytes([POLICY]) == (0, 0, 7 * 2) + + +def test_base_model_groups_own_no_adapter_gradients(monkeypatch): + r = rank() + values = r._estimate_flat_forward(requests(67, 4096)) + pending = [23 * 2**20] * 40 + [0] + with_pending(monkeypatch, r, pending) + n, out, signature, groups, head = values + values = dict( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=((33, True), (34, True), (4096, False)), + slot_refs=(POLICY, LoRASlotRef("checkpoint", None), None), + head_workspace_bytes=head, + ) + both = r._subforward_cost(**values) + # The base group's slot has no adapter, so a split with or without it + # still shares one slot's gradients. + assert both.checkpoint_adapter_gradient_slots == '[["checkpoint", "policy"]]' + # But its boundaries stay live while the policy group's backward runs. + boundary = 2048 * 2 * 40 + expected = sequential_oracle( + [(pending, [33 * 2048 * 2] * 40), ([0] * 41, [34 * 2048 * 2] * 40)] + ) + assert ( + both.checkpoint_adapter_gradient + == expected + == oracle(pending, [33 * 2048 * 2] * 40) + ) + assert expected > oracle(pending, [(33 + 34) * boundary // 40] * 40) + assert r._estimate_required_memory_bytes_from_values(**values) == both.required + + +def test_other_checkpoint_parameters_are_live_throughout(): + from art.trainer_rank._impl import _CheckpointSlot + + adapter = parameter(3) + custom = parameter(11) + r = adapter_rank([lora(policy=[adapter])], torch.nn.Module()) + r._checkpoint_slots["policy"] = _CheckpointSlot(params=(adapter, custom)) + # A custom object's parameter has no decoder position; LoRA ones keep theirs. + assert r._pending_adapter_gradient_bytes([POLICY]) == (3 * 2, 11 * 2) + + +def test_a_step_with_allocated_gradients_prices_no_extra(): + params = [parameter(100) for _ in range(4)] + r = adapter_rank([lora(policy=[p]) for p in params], torch.nn.Module()) + assert r._pending_adapter_gradient_bytes([POLICY]) == (200, 200, 200, 200, 0) + for p in params: + p.grad = torch.zeros_like(p) + # Later waves of the step find them in the availability baseline. + assert r._pending_adapter_gradient_bytes([POLICY]) == () + + +def test_slotless_lora_modules_hold_no_slot_parameters(): + module = LoRA.__new__(LoRA) + torch.nn.Module.__init__(module) + r = adapter_rank([module], torch.nn.Module()) + assert r._pending_adapter_gradient_bytes([POLICY]) == () + + +def priced(r, values, slot_refs): + n, out, signature, groups, head = values + return r._subforward_cost( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=slot_refs, + head_workspace_bytes=head, + ) + + +def test_cost_and_estimate_charge_the_extra_while_gradients_are_pending(monkeypatch): + r = rank() + values = r._estimate_flat_forward(requests(67, 4096)) + n, out, signature, groups, head = values + retained, workspace = r._checkpoint_memory_floor(groups) + boundary = retained // 40 + pending = [23 * 2**20] * 40 + [0] + with_pending(monkeypatch, r, pending) + extra = oracle(pending, [boundary] * 40) + assert extra > 0 + cost = priced(r, values, (POLICY, None)) + assert cost.checkpoint_adapter_gradient == extra + assert cost.checkpoint_adapter_gradient_slots == '[["checkpoint", "policy"]]' + assert cost.required == int( + (out + retained + workspace + COLD + retained + extra) * 1.1 + ) + estimate = r._estimate_required_memory_bytes_from_values( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=(POLICY, None), + head_workspace_bytes=head, + ) + assert estimate == cost.required + # Only gradient groups' slots count; no slot, no extra. + assert priced(r, values, (None, POLICY)).checkpoint_adapter_gradient == 0 + with_pending(monkeypatch, r, ()) + plain = priced(r, values, (POLICY, None)) + assert plain.checkpoint_adapter_gradient == 0 + assert plain.required == int((out + 2 * retained + workspace + COLD) * 1.1) + # A learned profile drops the first-execution transients, not the extra. + with_pending(monkeypatch, r, pending) + r._memory_profiles[signature] = _MemoryProfile(bytes_per_token=1, packed_tokens=n) + assert priced(r, values, (POLICY, None)).required == int( + (out + 2 * retained + workspace + extra) * 1.1 + ) + + +def child(extra: int, slots: str, workspace: int = 10) -> _SubforwardCost: + return _SubforwardCost( + required=int((1 + workspace + 1 + extra) * 1.1), + retained=0, + checkpoint_retained=1, + checkpoint_workspace=workspace, + checkpoint_input_gradient=1, + checkpoint_adapter_gradient=extra, + checkpoint_adapter_gradient_slots=slots, + ) + + +def test_split_charges_shared_gradients_once_and_distinct_slots_each(): + shared = [child(500, "a"), child(300, "a")] + assert TrainerRank._split_required_memory(shared) == int((2 + 2 + 10 + 500) * 1.1) + distinct = [child(500, "a"), child(300, "b")] + assert TrainerRank._split_required_memory(distinct) == int((2 + 2 + 10 + 800) * 1.1) + # A child with no pending gradients does not change the shared charge. + assert TrainerRank._split_required_memory([child(500, "a"), child(0, "")]) == int( + (2 + 2 + 10 + 500) * 1.1 + ) + + +def test_cheap_estimate_defers_while_a_gradient_slot_has_pending_gradients( + monkeypatch, +): + r = rank() + monkeypatch.setattr(r, "_ensure_checkpoint_slots_for", lambda *a, **k: None) + monkeypatch.setattr( + r, + "_resolve_slot_ref", + lambda request, checkpoint: POLICY if not request.no_grad else None, + ) + seen: list[tuple[LoRASlotRef, ...]] = [] + + def pending(refs): + seen.append(tuple(refs)) + return (1,) * 41 + + monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", pending) + assert r._estimate_flat_forward(requests(67, 4096)) is None + assert seen == [(POLICY,)] + monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", lambda refs: ()) + assert r._estimate_flat_forward(requests(67, 4096)) is not None + + +def test_a_base_only_gradient_group_prices_and_defers_nothing(monkeypatch): + r = rank() + values = r._estimate_flat_forward(requests(67, 4096)) + n, out, signature, groups, head = values + base = LoRASlotRef("checkpoint", None) + with_pending(monkeypatch, r, [23 * 2**20] * 40 + [0]) + cost = priced(r, values, (base, None)) + # The base model has no adapter: no extra, no slot identity. + assert cost.checkpoint_adapter_gradient == 0 + assert cost.checkpoint_adapter_gradient_slots == "" + estimate = r._estimate_required_memory_bytes_from_values( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=(base, None), + head_workspace_bytes=head, + ) + assert estimate == cost.required + # Nor does the cheap estimate defer to the exact plan for it. + monkeypatch.setattr(r, "_ensure_checkpoint_slots_for", lambda *a, **k: None) + monkeypatch.setattr(r, "_resolve_slot_ref", lambda request, checkpoint: base) + seen: list[tuple[LoRASlotRef, ...]] = [] + + def pending(refs): + seen.append(tuple(refs)) + return (1,) * 41 + + monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", pending) + assert r._estimate_flat_forward(requests(67, 4096)) is not None + assert seen == [] diff --git a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py index d6d9ded56..83c1ad50a 100644 --- a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py @@ -9,6 +9,9 @@ import torch from art.trainer_rank import ForwardInput +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) from art.trainer_rank._impl import Unset, _MemoryProfile, _SplitForwardPlan @@ -29,12 +32,15 @@ def test_pending_cold_peak_does_not_become_forward_retention(pending_rank): assert cost.checkpoint_input_gradient == gradient # Exact previous cold estimate, including outputs and its one safety factor. assert cost.retained == 23102959299 - assert cost.required == int((plan.output_bytes + 2 * gradient + 12705630112) * 1.1) + # Unprofiled: the first execution's transients beside the workspace. + assert cost.required == int( + (plan.output_bytes + 2 * gradient + 12705630112 + COLD) * 1.1 + ) assert r._memory_check(plan).estimated_required_bytes == cost.required profile(r, plan) warm = r._plan_cost(plan) assert warm.retained == int((plan.output_bytes + gradient) * 1.1) - assert warm.required == cost.required + assert warm.required == int((plan.output_bytes + 2 * gradient + 12705630112) * 1.1) @pytest.mark.parametrize("rows", [1, 67, 1024]) @@ -60,8 +66,8 @@ def test_gradient_is_not_absorbed_by_larger_head_workspace(): head = 10**10 cost = price(r, (n, out, sig, groups, head)) gradient = 67 * 40 * 2048 * 2 - assert cost.checkpoint_workspace == head - assert cost.required == int((out + head + 2 * gradient) * 1.1) + assert cost.checkpoint_workspace == head + COLD + assert cost.required == int((out + head + COLD + 2 * gradient) * 1.1) assert cost.retained == int((out + head + gradient) * 1.1) r._memory_profiles[sig] = _MemoryProfile( bytes_per_token=10**9, @@ -199,16 +205,18 @@ def test_split_priority_subtracts_only_uncovered_gradient_peak(fully_masked): r = rank() plan = r._plan_flat_forward(requests(17, 19)) cold = r._plan_cost(plan) + # Once profiled, the static estimate has no first-execution transients. + profile(r, plan) + static = r._plan_cost(plan).required + assert static < cold.required # Place a real learned peak between the two static estimates, or above both. - measured = ( - cold.required + 10**7 if fully_masked else (cold.retained + cold.required) / 2 - ) + measured = static + 10**7 if fully_masked else (cold.retained + static) / 2 rate = (measured / 1.1 - plan.output_bytes) / plan.packed_tokens profile(r, plan, rate=rate) cost = r._plan_cost(plan) old_required = int((plan.output_bytes + int(plan.packed_tokens * rate)) * 1.1) assert cold.retained < old_required - assert cost.required == max(cold.required, old_required) + assert cost.required == max(static, old_required) assert ( cost.ephemeral - cost.checkpoint_peak_increment == old_required - cost.retained ) diff --git a/tests/unit/test_trainer_rank_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py index 21b7ad071..54d279f91 100644 --- a/tests/unit/test_trainer_rank_head_memory.py +++ b/tests/unit/test_trainer_rank_head_memory.py @@ -8,6 +8,9 @@ import torch from art.trainer_rank import ForwardInput +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) from art.trainer_rank._impl import ( _PACKED_PRICED_LOGICAL_ROW_BYTES, Unset, @@ -147,7 +150,10 @@ def test_outputs_retention_and_empirical_peak_are_counted_once(): head = 3 * 512 * 248320 * 2 cost = r._plan_cost(plan) assert cost.retained == int((plan.output_bytes + retained + head) * 1.1) - assert cost.required == int((plan.output_bytes + retained + gradient + head) * 1.1) + # Unprofiled: the first execution's transients beside the head workspace. + assert cost.required == int( + (plan.output_bytes + retained + gradient + head + COLD) * 1.1 + ) r._memory_profiles[plan.signature] = _MemoryProfile( bytes_per_token=2_000_000, packed_tokens=512, @@ -312,8 +318,8 @@ def test_target_backward_refuses_budget_below_logits_and_both_gradients(rows): retained, _ = r._checkpoint_memory_floor(r._plan_group_rows(plan)) gradient = rows * 40 * 2048 * 2 dense = min(rows, 512) * 248320 * 2 - before = int((plan.output_bytes + retained + gradient + 2 * dense) * 1.1) - expected = int((plan.output_bytes + retained + gradient + 3 * dense) * 1.1) + before = int((plan.output_bytes + retained + gradient + 2 * dense + COLD) * 1.1) + expected = int((plan.output_bytes + retained + gradient + 3 * dense + COLD) * 1.1) r._available_memory_bytes = lambda: (before + expected) // 2 check = r._memory_check(plan) assert check.estimated_required_bytes == expected diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 3fbdcee89..9cb0315ba 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -9,6 +9,9 @@ import torch from art.trainer_rank import ForwardInput, TrainerRank +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) from art.trainer_rank._impl import ( _PACKED_PRICED_LOGICAL_ROW_BYTES, _MemoryProfile, @@ -741,15 +744,19 @@ def test_hybridep_recompute_prices_fresh_dense_output_without_buffer_growth( retained, workspace = rank._checkpoint_memory_floor(groups) assert workspace == 218752 * 2048 * 2 == 896008192 cost = rank._subforward_cost(**values) - assert cost.required == int((8 + 2 * retained + workspace) * 1.1) - assert cost.checkpoint_workspace == workspace # Maximum, not stage + output. + # Unprofiled: the first execution's transients sit beside the extent. + assert cost.required == int((8 + 2 * retained + workspace + COLD) * 1.1) + assert cost.checkpoint_workspace == workspace + COLD # Not stage + output. rank._available_memory_bytes = lambda: 600000000 assert rank._memory_check_required(baseline.required).fits assert not rank._memory_check_required(cost.required).fits rank._memory_profiles[signature] = _MemoryProfile( bytes_per_token=1, packed_tokens=2 ) - assert rank._subforward_cost(**values).required == cost.required + # Profiled: the floor alone, without first-execution transients. + assert rank._subforward_cost(**values).required == int( + (8 + 2 * retained + workspace) * 1.1 + ) rank._memory_profiles[signature] = _MemoryProfile( bytes_per_token=10**9, packed_tokens=2 ) diff --git a/tests/unit/test_trainer_rank_pending_memory.py b/tests/unit/test_trainer_rank_pending_memory.py index 359de8daa..8cae0619a 100644 --- a/tests/unit/test_trainer_rank_pending_memory.py +++ b/tests/unit/test_trainer_rank_pending_memory.py @@ -12,6 +12,7 @@ from art.megatron.prefix_tree_packing import prefix_tree_pack from art.trainer_rank import ForwardInput, TrainerRank from art.trainer_rank import _gdn_memory as g +from art.trainer_rank._impl import _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD from art.trainer_rank._impl import Unset, _MemoryProfile @@ -133,7 +134,7 @@ def test_actual_constructor_cache_and_full_plan(pending_rank): assert ( rank._memory_check(plan).estimated_required_bytes == rank._plan_cost(plan).required - == 32229502659 + == 32303322409 ) selected = rank._select_next_micro_batch(requests, 0) assert ( @@ -196,7 +197,7 @@ def test_exact_pending_demand_survives_recovery(monkeypatch, pending_rank, fits_ plan = pending_rank._plan_flat_forward(requests) assert pending_rank._estimate_flat_forward(requests) is None assert g.plan_floor(pending_rank, plan) == (8296857600, 12705630112) - assert pending_rank._memory_check(plan).estimated_required_bytes == 32229502659 + assert pending_rank._memory_check(plan).estimated_required_bytes == 32303322409 _check_component_demand_recovery( monkeypatch, pending_rank, requests, fits_after=fits_after ) @@ -214,8 +215,8 @@ def test_original_installed_norm_preserves_pending_floor(layer): assert g.model_shapes(rank) is not None plan = rank._plan_flat_forward(full_requests()) assert g.plan_floor(rank, plan) == (8296857600, 12705630112) - assert rank._memory_check(plan).estimated_required_bytes == 32229502659 - assert rank._plan_cost(plan).required == 32229502659 + assert rank._memory_check(plan).estimated_required_bytes == 32303322409 + assert rank._plan_cost(plan).required == 32303322409 assert rank._estimate_flat_forward(full_requests()) is None for requests in ([], full_requests(no_grad=True)): assert g.plan_floor(rank, rank._plan_flat_forward(requests)) == (0, 0) @@ -384,7 +385,7 @@ def test_constructor_declined_moe_keeps_generic_admission(layer, unsupported): required = rank._plan_cost(plan).required # Generic checkpoint-input accounting still applies without a MoE component. gradient = 50640 * 40 * 2048 * 2 - assert required == int((plan.output_bytes + 2 * gradient) * 1.1) + assert required == int((plan.output_bytes + 2 * gradient + COLD) * 1.1) rank._available_memory_bytes = lambda: required - 1 assert not rank._memory_check(plan).fits rank._available_memory_bytes = lambda: required diff --git a/tests/unit/test_trainer_rank_planner_reports.py b/tests/unit/test_trainer_rank_planner_reports.py index a58adfbc1..a9b26b6af 100644 --- a/tests/unit/test_trainer_rank_planner_reports.py +++ b/tests/unit/test_trainer_rank_planner_reports.py @@ -299,6 +299,8 @@ def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): "checkpoint_input_gradient": 0, "checkpoint_peak_increment": 0, "hybridep_growth": 0, + "checkpoint_adapter_gradient": 0, + "checkpoint_adapter_gradient_slots": "", }, } ], @@ -387,6 +389,7 @@ def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): "checkpoint_workspace", "checkpoint_peak_increment", "hybridep_growth", + "checkpoint_adapter_gradient", ): altered = json.loads(path.read_bytes()) altered["replay"]["memory_replay"]["estimates"][0]["cost_components"][ diff --git a/tests/unit/test_trainer_rank_shared_memory.py b/tests/unit/test_trainer_rank_shared_memory.py index 36bc971d7..d8b1a2938 100644 --- a/tests/unit/test_trainer_rank_shared_memory.py +++ b/tests/unit/test_trainer_rank_shared_memory.py @@ -126,7 +126,7 @@ def test_shared_return_in_actual_constructor_and_plan(layer, gate, no_grad): 8296857600, 50640 * (checkpoint_coefficient + 128) + 3157761952, ) - assert rank._plan_cost(plan).required == (32685829827 if gate else 32457666243) + assert rank._plan_cost(plan).required == (32759649577 if gate else 32531485993) selected = rank._select_next_micro_batch(requests, 0) assert ( selected.check.estimated_required_bytes @@ -147,7 +147,7 @@ def test_original_norm_installation_preserves_shared_return(layer, gated): 8296857600, 50640 * (checkpoint_coefficient + 128) + 3157761952, ) - expected = 32685829827 if gated else 32457666243 + expected = 32759649577 if gated else 32531485993 assert rank._memory_check(plan).estimated_required_bytes == expected assert rank._plan_cost(plan).required == expected diff --git a/tests/unit/test_trainer_rank_tp_floor.py b/tests/unit/test_trainer_rank_tp_floor.py index 6d5f4d00d..c3f49e3b9 100644 --- a/tests/unit/test_trainer_rank_tp_floor.py +++ b/tests/unit/test_trainer_rank_tp_floor.py @@ -14,6 +14,7 @@ import torch from art.trainer_rank import TrainerRank +from art.trainer_rank._impl import _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD from art.trainer_rank._impl import _MemorySignature H, F, LAYERS = 5120, 17408, 64 @@ -114,10 +115,11 @@ def test_the_traced_tp4_wave_prices_its_boundary_shards_and_their_repeat(): assert workspace == 3 * SEGMENT cost = _required(r) assert cost.checkpoint_input_gradient == retained - # One segment plus up to three TP-padding roots, each with its states. - state = cost.checkpoint_workspace + # One segment plus up to three TP-padding roots, each with its states, + # beside the unprofiled first execution's transients. + state = cost.checkpoint_workspace - COLD assert state == 4 * SEGMENT - assert cost.required == int((OUTPUT + 2 * retained + state) * 1.1) + assert cost.required == int((OUTPUT + 2 * retained + state + COLD) * 1.1) # Measured cold on all four ranks: 7.130 GB (7.060 GB in production), all # but the boundaries a transient recompute workspace; this raw floor # (8.43 GB) covers it. Today's cold admission was 4.637 GB. @@ -216,9 +218,9 @@ def test_gdn_segment_states_are_priced_with_the_segments(): r = tp_rank() rows = 8192 cost = _required(r, group_rows=((rows, True),), gdn_segments=4096) - assert cost.checkpoint_workspace == (4096 + 3) * SEGMENT + assert cost.checkpoint_workspace == (4096 + 3) * SEGMENT + COLD assert cost.required == int( - (OUTPUT + 2 * rows // 4 * LAYERS * H * 2 + (4096 + 3) * SEGMENT) * 1.1 + (OUTPUT + 2 * rows // 4 * LAYERS * H * 2 + (4096 + 3) * SEGMENT + COLD) * 1.1 ) @@ -227,7 +229,7 @@ def test_tp_padding_roots_carry_their_own_states(): r = tp_rank() cost = _required(r, group_rows=((4, True),), gdn_segments=1) # Four roots' initial states alone: 4 x 12 value heads x 128 x 128 x fp32. - assert cost.checkpoint_workspace == 4 * SEGMENT > 4 * 12 * 128 * 128 * 4 + assert cost.checkpoint_workspace - COLD == 4 * SEGMENT > 4 * 12 * 128 * 128 * 4 assert cost.required > 4 * SEGMENT