From f0670b9ec637b8143f55f3a4a2b4873ca92f2306 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 01:50:04 +0000 Subject: [PATCH 01/14] Price the adapter gradients a recompute backward holds at its peak A full-recompute backward allocates each layer's adapter gradients as it passes, and the step's optimizer frees them. Recomputing layer i still holds the saved boundaries of layers 0..i, so a short first wave peaks at layer 0 with nearly every layer's gradients live; the checkpoint floor priced only the last layer's end (all boundaries). On Qwen3.6-35B-A3B CP2, 2k and 4k token first waves were admitted 11.5% and 2.0% under their peaks. The floor now adds the largest excess of pending gradients over released boundaries across the real layers, while a slot's gradients are unallocated, and an unprofiled wave's 64 MiB of first-execution transients. Co-Authored-By: Claude Opus 5.5 (1M context) --- .github/workflows/prek.yml | 2 + src/art/trainer_rank/_impl.py | 170 +++++++++++- ...st_trainer_rank_adapter_gradient_memory.py | 258 ++++++++++++++++++ ...trainer_rank_checkpoint_gradient_memory.py | 24 +- tests/unit/test_trainer_rank_head_memory.py | 12 +- tests/unit/test_trainer_rank_moe_memory.py | 13 +- .../unit/test_trainer_rank_pending_memory.py | 11 +- .../unit/test_trainer_rank_planner_reports.py | 3 + tests/unit/test_trainer_rank_shared_memory.py | 4 +- tests/unit/test_trainer_rank_tp_floor.py | 14 +- 10 files changed, 482 insertions(+), 29 deletions(-) create mode 100644 tests/unit/test_trainer_rank_adapter_gradient_memory.py diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 18a49a302..67a05d36c 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 \ @@ -279,6 +280,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 d4f7bde0f..01d2953c5 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -125,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 @@ -1061,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`` identifies those slots within + # this process (a hash, 0 when none), keeping the cost JSON-serializable. + checkpoint_adapter_gradient: int = 0 + checkpoint_adapter_gradient_slots: int = 0 @property def ephemeral(self) -> int: @@ -3408,12 +3419,32 @@ def _split_required_memory(costs: Sequence[_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 * _MEMORY_SAFETY_FACTOR)) @@ -4213,6 +4244,88 @@ def _checkpoint_memory_floor( workspace = max(workspace, -(-rows // 4) * 4 * self._hidden_size * 2) return retained, workspace + def _pending_adapter_gradient_bytes( + self, 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 = _language_model(chunk).decoder.layers + except (AttributeError, RuntimeError): + return () + layer_of: dict[int, int] = {} + params: dict[int, torch.nn.Parameter] = {} + + def slot_params(module: torch.nn.Module) -> Iterator[torch.nn.Parameter]: + for child in module.modules(): + # A LoRA without slot tables holds no slot parameters. + if isinstance(child, LoRA) and "_slot_keys" in vars(child): + for ref in refs: + yield from child.lora_slot_params(ref) + + for index, layer in enumerate(layers): + for param in slot_params(layer): + params[id(param)] = param + layer_of[id(param)] = max(layer_of.get(id(param), -1), index) + for param in slot_params(chunk): + 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_adapter_gradient_bytes( + self, slots: Iterable["LoRASlotRef"], boundaries: Sequence[int] + ) -> int: + """The recompute backward's adapter-gradient peak beyond released boundaries. + + 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 the extra peak is + max(0, max over i of gradients(i..) - boundaries(i+1..)). It is taken + over the real per-layer sizes, not 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. + ``boundaries`` gives each decoder layer's saved-boundary bytes; + ``slots`` are the gradient groups' adapter slots. + """ + pending = self._pending_adapter_gradient_bytes(slots) + if not pending or len(pending) != len(boundaries) + 1: + return 0 + 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] + return extra + def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( packed_tokens=plan.packed_tokens, @@ -4275,18 +4388,39 @@ 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 = 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 + ) + adapter_gradient = ( + self._checkpoint_adapter_gradient_bytes( + gradient_slots, (gradient // self._num_layers,) * self._num_layers + ) + 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 ), ) @@ -4300,6 +4434,10 @@ 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=hash(gradient_slots) + if adapter_gradient + else 0, ) def _retained_memory_bytes( @@ -6060,6 +6198,17 @@ 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 + ] + 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 _gdn_memory.model_shapes(self) is not None @@ -8071,11 +8220,28 @@ def _estimate_required_memory_bytes_from_values( retained, workspace = self._checkpoint_memory_floor( group_rows, slot_refs, gdn_segments ) + 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( + ( + ref + for (_, grad), ref in zip( + group_rows, + slot_refs or (None,) * len(group_rows), + strict=True, + ) + if grad and ref is not None + ), + (retained // self._num_layers,) * self._num_layers, + ) + if profiled is None and any(grad for _, grad in group_rows): + backward += _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/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..dad274d7d --- /dev/null +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -0,0 +1,258 @@ +"""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 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 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_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 == hash(frozenset({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: int, 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, 7), child(300, 7)] + assert TrainerRank._split_required_memory(shared) == int((2 + 2 + 10 + 500) * 1.1) + distinct = [child(500, 7), child(300, 9)] + 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, 7), child(0, 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 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 1adce5472..22bddb218 100644 --- a/tests/unit/test_trainer_rank_planner_reports.py +++ b/tests/unit/test_trainer_rank_planner_reports.py @@ -298,6 +298,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": 0, }, } ], @@ -379,6 +381,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 From ecf2ce21bc6fcbb657cbcf9cd6044dfdb1f63821 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 02:35:34 +0000 Subject: [PATCH 02/14] Count head-tied and custom checkpoint gradients as live throughout Name gradient slots with sorted kind/name JSON instead of a hash. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 53 +++++++++++++------ ...st_trainer_rank_adapter_gradient_memory.py | 29 ++++++++-- .../unit/test_trainer_rank_planner_reports.py | 2 +- 3 files changed, 62 insertions(+), 22 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 01d2953c5..78dfbc76b 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -20,6 +20,7 @@ from dataclasses import field as dataclass_field from functools import partial import hashlib +import json import logging import math import os @@ -1068,10 +1069,10 @@ class _SubforwardCost: # 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`` identifies those slots within - # this process (a hash, 0 when none), keeping the cost JSON-serializable. + # 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: int = 0 + checkpoint_adapter_gradient_slots: str = "" @property def ephemeral(self) -> int: @@ -4273,21 +4274,38 @@ def _pending_adapter_gradient_bytes( layer_of: dict[int, int] = {} params: dict[int, torch.nn.Parameter] = {} - def slot_params(module: torch.nn.Module) -> Iterator[torch.nn.Parameter]: - for child in module.modules(): + 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(child, LoRA) and "_slot_keys" in vars(child): + if isinstance(module, LoRA) and "_slot_keys" in vars(module): for ref in refs: - yield from child.lora_slot_params(ref) + yield from module.lora_slot_params(ref) for index, layer in enumerate(layers): - for param in slot_params(layer): + for param in slot_params(layer.modules()): params[id(param)] = param layer_of[id(param)] = max(layer_of.get(id(param), -1), index) - for param in slot_params(chunk): - if id(param) not in params: - params[id(param)] = param - layer_of[id(param)] = len(layers) + # Outside the decoder (a head runs its backward first) is live + # throughout, even for a parameter a decoder layer also uses. + inside = {id(module) for module in layers.modules()} + outside = (module for module in chunk.modules() if id(module) not in inside) + 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 ( @@ -4308,8 +4326,9 @@ def _checkpoint_adapter_gradient_bytes( 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 the extra peak is - max(0, max over i of gradients(i..) - boundaries(i+1..)). It is taken - over the real per-layer sizes, not a uniform-layer line. A short + 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. ``boundaries`` gives each decoder layer's saved-boundary bytes; @@ -4435,9 +4454,11 @@ def _subforward_cost( checkpoint_peak_increment=required - forward_required, hybridep_growth=hybridep_growth_bytes, checkpoint_adapter_gradient=adapter_gradient, - checkpoint_adapter_gradient_slots=hash(gradient_slots) + checkpoint_adapter_gradient_slots=json.dumps( + sorted([ref.kind, ref.name] for ref in gradient_slots) + ) if adapter_gradient - else 0, + else "", ) def _retained_memory_bytes( diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py index dad274d7d..34b686ccc 100644 --- a/tests/unit/test_trainer_rank_adapter_gradient_memory.py +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -142,6 +142,25 @@ def test_pending_gradients_count_local_unallocated_slot_parameters(): 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_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()) @@ -184,7 +203,7 @@ def test_cost_and_estimate_charge_the_extra_while_gradients_are_pending(monkeypa assert extra > 0 cost = priced(r, values, (POLICY, None)) assert cost.checkpoint_adapter_gradient == extra - assert cost.checkpoint_adapter_gradient_slots == hash(frozenset({POLICY})) + assert cost.checkpoint_adapter_gradient_slots == '[["checkpoint", "policy"]]' assert cost.required == int( (out + retained + workspace + COLD + retained + extra) * 1.1 ) @@ -212,7 +231,7 @@ def test_cost_and_estimate_charge_the_extra_while_gradients_are_pending(monkeypa ) -def child(extra: int, slots: int, workspace: int = 10) -> _SubforwardCost: +def child(extra: int, slots: str, workspace: int = 10) -> _SubforwardCost: return _SubforwardCost( required=int((1 + workspace + 1 + extra) * 1.1), retained=0, @@ -225,12 +244,12 @@ def child(extra: int, slots: int, workspace: int = 10) -> _SubforwardCost: def test_split_charges_shared_gradients_once_and_distinct_slots_each(): - shared = [child(500, 7), child(300, 7)] + shared = [child(500, "a"), child(300, "a")] assert TrainerRank._split_required_memory(shared) == int((2 + 2 + 10 + 500) * 1.1) - distinct = [child(500, 7), child(300, 9)] + 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, 7), child(0, 0)]) == int( + assert TrainerRank._split_required_memory([child(500, "a"), child(0, "")]) == int( (2 + 2 + 10 + 500) * 1.1 ) diff --git a/tests/unit/test_trainer_rank_planner_reports.py b/tests/unit/test_trainer_rank_planner_reports.py index 22bddb218..b5677bf28 100644 --- a/tests/unit/test_trainer_rank_planner_reports.py +++ b/tests/unit/test_trainer_rank_planner_reports.py @@ -299,7 +299,7 @@ def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): "checkpoint_peak_increment": 0, "hybridep_growth": 0, "checkpoint_adapter_gradient": 0, - "checkpoint_adapter_gradient_slots": 0, + "checkpoint_adapter_gradient_slots": "", }, } ], From 067048eb5e5549239b5c74106189992f7bf5bfba Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 02:54:38 +0000 Subject: [PATCH 03/14] Order slot names safely and leave base-model groups out of adapter slots Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 63 ++++++++++++------- ...st_trainer_rank_adapter_gradient_memory.py | 25 ++++++++ 2 files changed, 67 insertions(+), 21 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 78dfbc76b..073991611 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4245,6 +4245,20 @@ def _checkpoint_memory_floor( workspace = max(workspace, -(-rows // 4) * 4 * self._hidden_size * 2) return retained, workspace + @staticmethod + 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, refs: Iterable["LoRASlotRef"] ) -> tuple[int, ...]: @@ -4288,9 +4302,18 @@ def slot_params( 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 a decoder layer also uses. - inside = {id(module) for module in layers.modules()} - outside = (module for module in chunk.modules() if id(module) not in inside) + # 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) @@ -4407,13 +4430,7 @@ 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 = 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 - ) + gradient_slots = self._gradient_slots(group_rows, slot_refs) adapter_gradient = ( self._checkpoint_adapter_gradient_bytes( gradient_slots, (gradient // self._num_layers,) * self._num_layers @@ -4455,7 +4472,17 @@ def _subforward_cost( hybridep_growth=hybridep_growth_bytes, checkpoint_adapter_gradient=adapter_gradient, checkpoint_adapter_gradient_slots=json.dumps( - sorted([ref.kind, ref.name] for ref in gradient_slots) + [ + [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 "", @@ -6220,7 +6247,9 @@ def _estimate_flat_forward( # 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 + ref + for (ref, grad), _ in groups + if grad and ref is not None and ref.name is not None ] if ( gradient_slots @@ -8245,15 +8274,7 @@ def _estimate_required_memory_bytes_from_values( 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( - ( - ref - for (_, grad), ref in zip( - group_rows, - slot_refs or (None,) * len(group_rows), - strict=True, - ) - if grad and ref is not None - ), + self._gradient_slots(group_rows, slot_refs), (retained // self._num_layers,) * self._num_layers, ) if profiled is None and any(grad for _, grad in group_rows): diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py index 34b686ccc..c3aa0ae79 100644 --- a/tests/unit/test_trainer_rank_adapter_gradient_memory.py +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -150,6 +150,31 @@ def test_a_parameter_the_head_also_uses_is_live_throughout(): 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)) + with_pending(monkeypatch, r, [23 * 2**20] * 40 + [0]) + n, out, signature, groups, head = values + both = r._subforward_cost( + 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, + ) + # 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"]]' + + def test_other_checkpoint_parameters_are_live_throughout(): from art.trainer_rank._impl import _CheckpointSlot From c35b238798076e6c34e02d87823b32cc74184e01 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 06:33:29 +0000 Subject: [PATCH 04/14] Test that base-model gradient groups price and defer nothing Co-Authored-By: Claude Opus 5.5 (1M context) --- ...st_trainer_rank_adapter_gradient_memory.py | 34 +++++++++++++++++++ 1 file changed, 34 insertions(+) diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py index c3aa0ae79..627499d3a 100644 --- a/tests/unit/test_trainer_rank_adapter_gradient_memory.py +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -300,3 +300,37 @@ def pending(refs): 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 == [] From e54f4bc25814f9d547145721b5695661d5308c54 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:08:04 +0000 Subject: [PATCH 05/14] Walk gradient groups' backward one after another, in the worst order Autograd drains the last-forwarded group's chain before an earlier one's, and separate backward calls may come in either order, so a short group's gradients can peak beside another group's unreleased boundaries. Price each gradient group against its own boundaries, with groups not yet run holding theirs and groups already run holding their gradients, over every order. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 73 ++++++++--- ...st_trainer_rank_adapter_gradient_memory.py | 114 +++++++++++++++++- 2 files changed, 167 insertions(+), 20 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 073991611..46750c08d 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -20,6 +20,7 @@ from dataclasses import field as dataclass_field from functools import partial import hashlib +import itertools import json import logging import math @@ -4339,8 +4340,31 @@ def slot_params( sizes[layer_of[param_id]] += param.numel() * param.element_size() return tuple(sizes) if any(sizes) else () + def _checkpoint_gradient_groups( + self, + 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, slots: Iterable["LoRASlotRef"], boundaries: Sequence[int] + self, groups: Sequence[tuple["LoRASlotRef | None", Sequence[int]]] ) -> int: """The recompute backward's adapter-gradient peak beyond released boundaries. @@ -4354,19 +4378,39 @@ def _checkpoint_adapter_gradient_bytes( 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. - ``boundaries`` gives each decoder layer's saved-boundary bytes; - ``slots`` are the gradient groups' adapter slots. + ``groups`` gives each gradient group's adapter slot (None for the base + model) and each decoder layer's saved-boundary bytes. 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, a group yet to + run still holds all its boundaries and one already run all its + gradients, so take the worst order. """ - pending = self._pending_adapter_gradient_bytes(slots) - if not pending or len(pending) != len(boundaries) + 1: + 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)) + if not any(any(pending) for pending, _ in chains): return 0 - 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] - return extra + if len(chains) > 4: + # Too many orders to walk: every gradient live, nothing released. + return sum(sum(pending) for pending, _ in chains) + worst = 0 + for order in itertools.permutations(chains): + allocated = released = 0 + for pending, boundaries in order: + gradients = allocated + pending[-1] + worst = max(worst, gradients - released) + for index in range(len(boundaries) - 1, -1, -1): + gradients += pending[index] + worst = max(worst, gradients - released) + released += boundaries[index] + allocated = gradients + return worst def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( @@ -4433,7 +4477,7 @@ def _subforward_cost( gradient_slots = self._gradient_slots(group_rows, slot_refs) adapter_gradient = ( self._checkpoint_adapter_gradient_bytes( - gradient_slots, (gradient // self._num_layers,) * self._num_layers + self._checkpoint_gradient_groups(group_rows, slot_refs) ) if gradient else 0 @@ -8274,8 +8318,7 @@ def _estimate_required_memory_bytes_from_values( 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._gradient_slots(group_rows, slot_refs), - (retained // self._num_layers,) * self._num_layers, + self._checkpoint_gradient_groups(group_rows, slot_refs) ) if profiled is None and any(grad for _, grad in group_rows): backward += _COLD_RECOMPUTE_TRANSIENT_BYTES diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py index 627499d3a..0e80b9e2e 100644 --- a/tests/unit/test_trainer_rank_adapter_gradient_memory.py +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -5,6 +5,7 @@ """ from collections.abc import Sequence +import itertools import random import pytest @@ -61,7 +62,7 @@ def test_extra_is_the_worst_backward_layer_over_real_sizes( ): r = rank() with_pending(monkeypatch, r, pending) - assert r._checkpoint_adapter_gradient_bytes((POLICY,), boundaries) == oracle( + assert r._checkpoint_adapter_gradient_bytes(((POLICY, boundaries),)) == oracle( pending, boundaries ) @@ -76,14 +77,103 @@ def test_extra_matches_every_layer_for_random_sizes(monkeypatch): ] boundaries = [generator.randint(0, 40) for _ in range(layers)] with_pending(monkeypatch, r, pending) - extra = r._checkpoint_adapter_gradient_bytes((POLICY,), boundaries) + 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 + assert r._checkpoint_adapter_gradient_bytes(((POLICY, [4] * 40),)) == 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 = [POLICY, OTHER, LoRASlotRef("checkpoint", "third")] + for _ in range(200): + layers = generator.randint(1, 6) + chains, groups, by_slot = [], [], {} + for slot in slots[: generator.randint(1, 3)]: + 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_every_gradient_live(monkeypatch): + r = rank() + slots = [LoRASlotRef("checkpoint", f"slot{index}") for index in range(5)] + pending = [3, 1, 2] + with_slot_pending(monkeypatch, r, {(slot,): pending for slot in slots}) + groups = [(slot, [100, 100]) for slot in slots] + # Too many orders to walk: all five slots' gradients, nothing released. + assert r._checkpoint_adapter_gradient_bytes(groups) == 5 * sum(pending) + assert r._checkpoint_adapter_gradient_bytes(groups[:4]) == sequential_oracle( + [(pending, [100, 100])] * 4 + ) def lora(**slots: list[torch.nn.Parameter]) -> LoRA: @@ -159,9 +249,10 @@ def test_a_module_the_head_also_uses_is_live_throughout(): def test_base_model_groups_own_no_adapter_gradients(monkeypatch): r = rank() values = r._estimate_flat_forward(requests(67, 4096)) - with_pending(monkeypatch, r, [23 * 2**20] * 40 + [0]) + pending = [23 * 2**20] * 40 + [0] + with_pending(monkeypatch, r, pending) n, out, signature, groups, head = values - both = r._subforward_cost( + values = dict( packed_tokens=n, output_bytes=out, signature=signature, @@ -170,9 +261,22 @@ def test_base_model_groups_own_no_adapter_gradients(monkeypatch): 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 0df629cf516a6237afb77a73fa60ea111c2b398c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:29:36 +0000 Subject: [PATCH 06/14] Price gradient groups' worst backward order in closed form Any set of the other groups can have run first, so the worst order adds every other group whose gradients outweigh its boundaries to one group's own walk. Exact for any number of groups, without walking permutations. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 72 ++++++++++--------- ...st_trainer_rank_adapter_gradient_memory.py | 23 +++--- 2 files changed, 52 insertions(+), 43 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 46750c08d..827f7ffa7 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -20,7 +20,6 @@ from dataclasses import field as dataclass_field from functools import partial import hashlib -import itertools import json import logging import math @@ -4368,23 +4367,9 @@ def _checkpoint_adapter_gradient_bytes( ) -> int: """The recompute backward's adapter-gradient peak beyond released boundaries. - 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 the 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`` gives each gradient group's adapter slot (None for the base - model) and each decoder layer's saved-boundary bytes. 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, a group yet to - run still holds all its boundaries and one already run all its - gradients, so take the worst order. + model) and each decoder layer's saved-boundary bytes + (``_adapter_gradient_walk``). """ chains = [] for slot, boundaries in groups: @@ -4394,22 +4379,45 @@ def _checkpoint_adapter_gradient_bytes( if pending and len(pending) != len(boundaries) + 1: return 0 chains.append((pending or (0,) * (len(boundaries) + 1), boundaries)) - if not any(any(pending) for pending, _ in chains): - return 0 - if len(chains) > 4: - # Too many orders to walk: every gradient live, nothing released. - return sum(sum(pending) for pending, _ in chains) + return self._adapter_gradient_walk(chains) + + @staticmethod + 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 order in itertools.permutations(chains): - allocated = released = 0 - for pending, boundaries in order: - gradients = allocated + pending[-1] - worst = max(worst, gradients - released) - for index in range(len(boundaries) - 1, -1, -1): - gradients += pending[index] - worst = max(worst, gradients - released) - released += boundaries[index] - allocated = gradients + 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 _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py index 0e80b9e2e..71721224a 100644 --- a/tests/unit/test_trainer_rank_adapter_gradient_memory.py +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -146,11 +146,11 @@ def test_gradient_groups_run_their_backward_one_after_another(monkeypatch): def test_sequential_groups_match_every_order_for_random_sizes(monkeypatch): r = rank() generator = random.Random(1) - slots = [POLICY, OTHER, LoRASlotRef("checkpoint", "third")] + 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, 3)]: + 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: @@ -163,17 +163,18 @@ def test_sequential_groups_match_every_order_for_random_sizes(monkeypatch): assert r._checkpoint_adapter_gradient_bytes(groups) == sequential_oracle(chains) -def test_many_gradient_groups_price_every_gradient_live(monkeypatch): +def test_many_gradient_groups_price_exactly(monkeypatch): r = rank() - slots = [LoRASlotRef("checkpoint", f"slot{index}") for index in range(5)] - pending = [3, 1, 2] - with_slot_pending(monkeypatch, r, {(slot,): pending for slot in slots}) - groups = [(slot, [100, 100]) for slot in slots] - # Too many orders to walk: all five slots' gradients, nothing released. - assert r._checkpoint_adapter_gradient_bytes(groups) == 5 * sum(pending) - assert r._checkpoint_adapter_gradient_bytes(groups[:4]) == sequential_oracle( - [(pending, [100, 100])] * 4 + 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: From 2a6419528d8f4071d03e80c2983bd062746345ee Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:44:22 +0000 Subject: [PATCH 07/14] Test that a layer-count mismatch prices no adapter-gradient extra Co-Authored-By: Claude Opus 5.5 (1M context) --- .../test_trainer_rank_adapter_gradient_memory.py | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py index 71721224a..0da30ce0b 100644 --- a/tests/unit/test_trainer_rank_adapter_gradient_memory.py +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -87,6 +87,20 @@ def test_no_pending_gradients_add_nothing(monkeypatch): 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. From be1f15df12409d156606457cb09ad39a5f13571c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 03:51:09 +0000 Subject: [PATCH 08/14] Replay the adapter-gradient floor from frozen slot facts Grouped planner replay (#1028) recomputes each subforward's cost from primitive runtime facts, with group indices standing in for slot refs. The adapter-gradient floor reads the gradient slot's unallocated LoRA gradients from the live model, so replay could neither resolve the slot nor reproduce the term. Capture each gradient group's slot kind and name and its pending gradient bytes per decoder layer with the selection (runtime facts version 2), and have ReplayRank answer the floor's slot and pending-gradient readers from those facts. The capture's stock-estimator check now covers the floor's readers too. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 92 ++++++++++++++++++- tests/unit/test_grouped_planner_replay.py | 55 ++++++++++- .../unit/test_planner_replay_owner_budget.py | 2 +- 3 files changed, 143 insertions(+), 6 deletions(-) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index f0da8bbe1..15786f0b0 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -11,7 +11,7 @@ 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 @@ -21,6 +21,12 @@ _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", @@ -96,12 +102,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 +225,18 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: if backwards else 0 ) + adapter = None + if group.grad_enabled and getattr(group.slot_ref, "name", None) 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 +254,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 +277,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 +314,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 +343,7 @@ def integer(value: Any, *, minimum: int = 0) -> None: "gradient", "head_rows", "head_target_rows", + "adapter", "gdn", }, ) @@ -363,6 +388,22 @@ 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) gdn = group["gdn"] if gdn is not None: fields(gdn, {"layers", "shapes", "segments"}) @@ -420,11 +461,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..5f63d0084 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -171,6 +171,59 @@ def test_selected_slot_terms_are_replayed_and_frozen(layer, tmp_path): ) +def test_selected_adapter_gradients_are_replayed_and_frozen(layer, tmp_path): + 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) + 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", ["version", "group", "layout", "gdn_segment", "budget"] ) @@ -184,7 +237,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) From b8a2fd183edabe1dae85a4e93e0e548238d2339f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 03:55:54 +0000 Subject: [PATCH 09/14] Narrow the replay capture's slot name before sizing it Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index 15786f0b0..92a746b66 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -226,7 +226,7 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: else 0 ) adapter = None - if group.grad_enabled and getattr(group.slot_ref, "name", None) is not 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) From 6b1b918aabaa3e3a0a2646057bcd0afe2b2aa083 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:12:59 +0000 Subject: [PATCH 10/14] Bound replay's recorded layer count and tighten adapter facts Replay sizes per-layer boundary tuples from the report's recorded num_layers, which only capture bounded; refuse counts outside capture's 1024-layer limit before any estimator runs. A slot without a kind (a megatron-less reference) has no pending gradients, so reject kindless facts that claim some. Test forged adapter facts, the layer bound, and an instance-overridden pending-gradient reader. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_misses.py | 3 ++ src/art/trainer_rank/_planner_replay.py | 5 ++- tests/unit/test_grouped_planner_replay.py | 55 ++++++++++++++++++++++- 3 files changed, 61 insertions(+), 2 deletions(-) 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 92a746b66..45c7acc22 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -17,6 +17,8 @@ _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 @@ -74,7 +76,7 @@ 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 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 @@ -393,6 +395,7 @@ def integer(value: Any, *, minimum: int = 0) -> None: fields(adapter, {"kind", "name", "pending"}) if ( not group["grad"] + or (adapter["kind"] is None and any(adapter["pending"] or ())) 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 diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 5f63d0084..29d384df6 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -171,7 +171,8 @@ def test_selected_slot_terms_are_replayed_and_frozen(layer, tmp_path): ) -def test_selected_adapter_gradients_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 @@ -197,6 +198,11 @@ def test_selected_adapter_gradients_are_replayed_and_frozen(layer, tmp_path): 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" @@ -224,6 +230,53 @@ def test_selected_adapter_gradients_are_replayed_and_frozen(layer, tmp_path): ) +@pytest.mark.parametrize( + "change", ["gradient", "length", "value", "kind_length", "kindless_pending"] +) +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 + else: + adapter["kind"] = None + with pytest.raises(ValueError): + 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 + 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"] ) From eedef3d110f153c30087b93670170d5b5eea753d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:16:20 +0000 Subject: [PATCH 11/14] Patch the reader with monkeypatch in the custom-reader test Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_grouped_planner_replay.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 29d384df6..ed3acc6c2 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -272,7 +272,9 @@ def test_custom_adapter_gradient_reader_is_explicitly_incomplete(monkeypatch): rank = _rank(monkeypatch) plan = rank._plan_flat_forward([_request(1)]) original = rank._pending_adapter_gradient_bytes - rank._pending_adapter_gradient_bytes = lambda refs: original(refs) + monkeypatch.setattr( + rank, "_pending_adapter_gradient_bytes", lambda refs: original(refs) + ) with pytest.raises(ValueError, match="custom_runtime_estimator"): _planner_replay.capture(rank, plan) From a4f05f82d860d9b9ca2914a2fae4d1dffab389e3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:27:30 +0000 Subject: [PATCH 12/14] Validate adapter pending facts before scanning them Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index 45c7acc22..739ee8360 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -395,7 +395,6 @@ def integer(value: Any, *, minimum: int = 0) -> None: fields(adapter, {"kind", "name", "pending"}) if ( not group["grad"] - or (adapter["kind"] is None and any(adapter["pending"] or ())) 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 @@ -407,6 +406,9 @@ def integer(value: Any, *, minimum: int = 0) -> None: 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"}) From 3fa5dbfd1093b47b3848cf3c1b2372fe5a122ab9 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:32:49 +0000 Subject: [PATCH 13/14] Refuse mixed slot kinds in facts and a zero layer count at capture Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 13 ++++++++++++- tests/unit/test_grouped_planner_replay.py | 11 +++++++++-- 2 files changed, 21 insertions(+), 3 deletions(-) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index 739ee8360..c3b7db458 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -76,7 +76,10 @@ def reserve(size: int) -> None: def capture(rank: Any, plan: Any) -> dict[str, Any]: - if rank._num_layers > MAX_LAYERS 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 @@ -436,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") diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index ed3acc6c2..5ee494755 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -231,7 +231,8 @@ def test_selected_adapter_gradients_are_replayed_and_frozen(layer, tmp_path): @pytest.mark.parametrize( - "change", ["gradient", "length", "value", "kind_length", "kindless_pending"] + "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) @@ -247,8 +248,14 @@ def test_adapter_fact_validation_rejects_forged_input(change, layer, tmp_path): adapter["pending"][0] = 1.5 elif change == "kind_length": adapter["kind"] = "k" * 65 - else: + 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": []} with pytest.raises(ValueError): reports.replay(report) From c41148ec8c0077aa1504970d643a833fbad7a07e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:44:24 +0000 Subject: [PATCH 14/14] Name the refusal each forged adapter fact must hit Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_grouped_planner_replay.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 5ee494755..c0c4b1ba8 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -256,7 +256,13 @@ def test_adapter_fact_validation_rejects_forged_input(change, layer, tmp_path): ] groups[0]["grad"] = True groups[0]["adapter"] = {"kind": None, "name": "base", "pending": []} - with pytest.raises(ValueError): + # 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)