From 8a7052de5abc6a01bbc929a0537ee42fc727b126 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 22 Sep 2026 02:33:59 +0000 Subject: [PATCH 1/9] Local baseline: exact #888 composition on c6ac48f3 --- .github/workflows/prek.yml | 4 + src/art/trainer_rank/_impl.py | 155 +++++++-- .../unit/test_trainer_rank_handoff_budget.py | 170 ++++++++++ .../test_trainer_rank_physical_reserve.py | 317 ++++++++++++++++++ ...trainer_rank_recovery_slots_distributed.py | 46 ++- tests/unit/test_trainer_rank_split_peak.py | 9 +- 6 files changed, 661 insertions(+), 40 deletions(-) create mode 100644 tests/unit/test_trainer_rank_handoff_budget.py create mode 100644 tests/unit/test_trainer_rank_physical_reserve.py diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 1f8f27e54..dde2dcf29 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -223,6 +223,8 @@ jobs: tests/unit/test_prefix_tree_attention_builder.py \ tests/unit/test_prefix_tree_grad_parity.py \ tests/unit/test_prefix_tree_packing.py \ + tests/unit/test_trainer_rank_handoff_budget.py \ + tests/unit/test_trainer_rank_physical_reserve.py \ tests/unit/test_trainer_rank_validation.py \ tests/unit/test_trainer_rank_weird_shapes.py \ tests/unit/test_trainer_rank_split.py \ @@ -252,5 +254,7 @@ jobs: --ignore=tests/unit/test_prefix_tree_attention_builder.py \ --ignore=tests/unit/test_prefix_tree_grad_parity.py \ --ignore=tests/unit/test_prefix_tree_packing.py \ + --ignore=tests/unit/test_trainer_rank_handoff_budget.py \ + --ignore=tests/unit/test_trainer_rank_physical_reserve.py \ --ignore=tests/unit/test_trainer_rank_validation.py \ --ignore=tests/unit/test_trainer_rank_weird_shapes.py diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index d08ce96e4..2825f3d45 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -2363,24 +2363,37 @@ def _forward_micro_batches( items, start, checkpoint=checkpoint ) self._snapshot_planning_telemetry(candidate.plan, candidate.check) - if isinstance(candidate.plan, _FlatForwardPlan): - tracked_outputs, memory_baseline = ( - self._run_flat_plan_with_memory_tracking( - candidate.plan, - check=candidate.check, - context="forward_micro_batches", + tracked_outputs: list[AnyForwardOutput] = [] + outputs: list[Any] = [] + flat_outputs = iter(tracked_outputs) + error: BaseException | None = None + try: + if isinstance(candidate.plan, _FlatForwardPlan): + tracked_outputs, memory_baseline = ( + self._run_flat_plan_with_memory_tracking( + candidate.plan, + check=candidate.check, + context="forward_micro_batches", + ) ) - ) - else: - tracked_outputs, memory_baseline, forward_peak = ( - self._execute_split_plan_with_memory_tracking( - candidate.plan, - check=candidate.check, - context="forward_micro_batches", + else: + tracked_outputs, memory_baseline, forward_peak = ( + self._execute_split_plan_with_memory_tracking( + candidate.plan, + check=candidate.check, + context="forward_micro_batches", + ) ) - ) - flat_outputs = iter(tracked_outputs) - outputs = [_unflatten(item, flat_outputs) for item in candidate.inputs] + flat_outputs = iter(tracked_outputs) + outputs = [_unflatten(item, flat_outputs) for item in candidate.inputs] + except BaseException as exc: + error = exc + try: + self._release_cached_memory_for_backward(candidate.plan, error=error) + except BaseException: + # Do not retain our completed graph through a new handoff traceback. + del tracked_outputs, flat_outputs, outputs + raise stop = start + candidate.stats_global_count if stop < len(items): self._last_global_micro_batch_size = max( @@ -2434,6 +2447,40 @@ def _forward_micro_batches( del tracked_outputs, flat_outputs, outputs start = stop + def _release_cached_memory_for_backward( + self, plan: _AnyForwardPlan, *, error: BaseException | None = None + ) -> None: + # Every WORLD wave reaches this before the public iterator skips empty + # outputs. Forward has already executed: never replan or retry here. + with self._cache_recovery_episode() as (owner, started): + exchange_error: BaseException | None = None + try: + failed, gradients = self._recovery_reduce( + [ + float(error is not None), + float(any(group.grad_enabled for group in plan.groups)), + ], + op="MAX", + sync_across_dp=True, + ) + except BaseException as exc: + if error is None: + raise + exchange_error = exc + if error is not None: + raise self._memory_error_with_reduction_note(error, exchange_error) + if failed: + raise RuntimeError("Forward failed on another rank before handoff") + if not gradients: + return + self._try_cache_recovery( + None, + sync_across_dp=True, + owner=owner, + started=started, + handoff_grad=any(group.grad_enabled for group in plan.groups), + ) + @overload def dp_rank_forward( self, @@ -5186,14 +5233,7 @@ def finish(value: Any) -> Any: return result assert refused is not None original = refused.error(context) - state = self._recovery_state() - started = self._recovery_clock() - owner = object() - with state.lock: - if state.owner is None: - state.owner = owner - primary: BaseException | None = None - try: + with self._cache_recovery_episode() as (owner, started): if not isinstance(value, _ForwardRefusal): # A formerly fitting width is not proof that the minimum cannot fit. value = search() @@ -5220,6 +5260,18 @@ def finish(value: Any) -> Any: self._snapshot_planning_telemetry(refused.plan, refused.check) latest = refused.error(context) raise latest from original + + @contextmanager + def _cache_recovery_episode(self) -> Iterator[tuple[object, float | None]]: + state = self._recovery_state() + started = self._recovery_clock() + owner = object() + with state.lock: + if state.owner is None: + state.owner = owner + primary: BaseException | None = None + try: + yield owner, started except BaseException as exc: primary = exc raise @@ -5334,11 +5386,12 @@ def _memory_error_with_reduction_note( def _try_cache_recovery( self, - check: _MemoryCheck, + check: _MemoryCheck | None, *, sync_across_dp: bool, owner: object, started: float | None, + handoff_grad: bool = False, ) -> bool: state = self._recovery_state() now = self._recovery_clock() @@ -5370,7 +5423,7 @@ def _try_cache_recovery( invalid |= any(not math.isfinite(value) for value in (*costs, sum(costs))) values = self._recovery_reduce( [ - float(check.estimated_required_bytes), + float(check.estimated_required_bytes) if check is not None else 0.0, 0.0 if invalid else state.work, float(state.first_consumed), float(invalid), @@ -5387,16 +5440,20 @@ def _try_cache_recovery( needed = False cap_blocks = False try: - available = self._available_memory_bytes() + available = self._available_memory_bytes() if check is not None else 0 if ( - available < required + (available < required if check is not None else handoff_grad) and self.device.type == "cuda" and torch.cuda.is_available() and torch.cuda.get_allocator_backend() == "native" ): free, total = torch.cuda.mem_get_info(self.device) needed = int(free) < required + int(total * _MEMORY_RESERVE_FRACTION) - if os.environ.get(_TEST_HOOKS_ENV) == "1": + if check is None: + needed &= int(torch.cuda.memory_reserved(self.device)) > int( + torch.cuda.memory_allocated(self.device) + ) + elif os.environ.get(_TEST_HOOKS_ENV) == "1": limit = os.environ.get(_TEST_MEMORY_LIMIT_ENV) if limit: cap_blocks = required > max( @@ -5426,7 +5483,7 @@ def _try_cache_recovery( if sampled[0] < 0: raise RuntimeError("Memory recovery sampling failed on another rank") state.invalid |= not bool(sampled[3]) - if required <= sampled[0]: + if check is not None and required <= sampled[0]: return True if sampled[1] == 0 or sampled[2] == 0 or state.invalid: return False @@ -5444,9 +5501,37 @@ def _try_cache_recovery( # physical condition again immediately before the sole call. free, total = torch.cuda.mem_get_info(self.device) if int(free) < required + int(total * _MEMORY_RESERVE_FRACTION): - attempted = True - torch.cuda.empty_cache() - available = self._available_memory_bytes() + if check is not None: + attempted = True + torch.cuda.empty_cache() + else: + allocated = int(torch.cuda.memory_allocated(self.device)) + reserved = int(torch.cuda.memory_reserved(self.device)) + if reserved > allocated: + # A soft trigger, not calibrated library demand. The + # native release affects unused caches process-wide. + evidence = dict( + device=str(self.device), + reserve_trigger_bytes=int( + total * _MEMORY_RESERVE_FRACTION + ), + physical_free_before_bytes=int(free), + allocated_bytes=allocated, + reserved_before_bytes=reserved, + ) + with _telemetry_phase( + "gradient_handoff_cache_release", evidence + ): + attempted = True + with torch.cuda.device(self.device): + torch.cuda.empty_cache() + evidence["physical_free_after_bytes"] = int( + torch.cuda.mem_get_info(self.device)[0] + ) + evidence["reserved_after_bytes"] = int( + torch.cuda.memory_reserved(self.device) + ) + available = self._available_memory_bytes() if check is not None else 0 except BaseException as exc: error, available = exc, -1 exchange_error: BaseException | None = None @@ -5466,7 +5551,9 @@ def _try_cache_recovery( raise self._memory_error_with_reduction_note(error, exchange_error) if sampled[0] < 0: raise RuntimeError("Memory recovery failed on another rank") - # Rebuild pure search caches even when this fresh sample decreased. + # Admission rebuilds pure search caches even when the sample decreased. + # Handoff ignores this return: denied/insufficient recovery still yields + # the completed outputs, with no claim that backward will fit. return True def _memory_check_required( diff --git a/tests/unit/test_trainer_rank_handoff_budget.py b/tests/unit/test_trainer_rank_handoff_budget.py new file mode 100644 index 000000000..50527bbb7 --- /dev/null +++ b/tests/unit/test_trainer_rank_handoff_budget.py @@ -0,0 +1,170 @@ +"""Shared admission/handoff budget and original iterator ownership, CPU only.""" + +from contextlib import nullcontext +from types import SimpleNamespace +import weakref + +import pytest +import torch + +from art.trainer_rank import ForwardOutput, TrainerRank, _impl +from tests.unit import test_trainer_rank_cache_recovery as recovery +from tests.unit.test_trainer_rank_physical_reserve import allocator +from tests.unit.test_trainer_rank_validation import ( + _runtime, + _stub_forward, + _target_request, +) + + +@pytest.fixture +def rig(): + fixture = recovery.TestRecovery() + rank, cuda, clock, ns = fixture.make() + cuda.memory_reserved = lambda device: cuda.allocated + 100 + cuda.device = lambda device: nullcontext() + try: + yield rank, cuda, clock, ns + finally: + fixture.doCleanups() + + +def plan(*modes): + return SimpleNamespace(groups=[SimpleNamespace(grad_enabled=m) for m in modes]) + + +def test_admission_consumes_the_only_first_release(rig): + rank, cuda, _, ns = rig + _, error, _ = recovery.run(rank, [recovery.fail(ns), recovery.success(ns)]) + assert error is None and cuda.events.count("release") == 1 + cuda.free = 1 + before = rank._recovery_state().cost + rank._release_cached_memory_for_backward(plan(True)) + assert cuda.events.count("release") == 1 + assert rank._recovery_state().cost > before + rank._record_recovery_work("forward_micro_batches", 100.0) + rank._release_cached_memory_for_backward(plan(True)) + assert cuda.events.count("release") == 2 + + +def test_handoff_consumes_the_only_first_release(rig): + rank, cuda, _, ns = rig + cuda.free = 1 + rank._release_cached_memory_for_backward(plan(True)) + assert cuda.events.count("release") == 1 + cuda.free = 40 + _, error, calls = recovery.run(rank, [recovery.fail(ns), recovery.success(ns)]) + assert isinstance(error, _impl.TrainerRankMemoryError) + assert calls == 1 and cuda.events.count("release") == 1 + + +@pytest.mark.parametrize("cost, releases", [(40.0, 1), (40.5, 0)]) +def test_repeat_budget_has_the_same_five_percent_boundary(rig, cost, releases): + rank, cuda, _, _ = rig + state = rank._recovery_state() + state.first_consumed, state.work, state.cost, state.high = True, 1000.0, cost, 10.0 + ticks = iter((1.0, 1.0, 2.0)) + rank._recovery_clock = lambda: next(ticks) + cuda.free = 1 + rank._release_cached_memory_for_backward(plan(True)) + assert cuda.events.count("release") == releases + assert state.cost == cost + 1.0 and state.work == 1000.0 and state.owner is None + + +@pytest.mark.parametrize("modes", [(), (False,)]) +def test_empty_and_local_no_grad_peer_participate_without_cuda_queries(rig, modes): + rank, cuda, _, _ = rig + calls = [] + + def reduce(values, *, op, sync_across_dp): + assert sync_across_dp + calls.append((op, len(values))) + if op == "MAX" and len(values) == 2: + values[1] = 1.0 # Only the other peer has local gradient work. + elif op == "MIN": + values[1] = -1.0 # The other peer needs/attempts a release. + return values + + rank._recovery_reduce = reduce + rank._release_cached_memory_for_backward(plan(*modes)) + assert calls == [("MAX", 2), ("SUM", 2), ("MAX", 4), ("MIN", 4), ("MIN", 2)] + assert cuda.events == [] and rank._recovery_state().first_consumed + + +@pytest.mark.parametrize("cancel", [False, True]) +def test_forward_error_survives_secondary_exchange_failure(rig, cancel): + rank, cuda, _, _ = rig + primary = KeyboardInterrupt("cancel") if cancel else RuntimeError("forward") + cause, context = ValueError("cause"), LookupError("context") + primary.__cause__ = cause + primary.__context__ = context + primary.__suppress_context__ = True + + def broken(values, **kwargs): + assert values[0] == 1.0 + raise OSError("secondary exchange") + + rank._recovery_reduce = broken + with pytest.raises(type(primary)) as caught: + rank._release_cached_memory_for_backward(plan(True), error=primary) + assert caught.value is primary and primary.__cause__ is cause + assert primary.__context__ is context and primary.__suppress_context__ + assert cuda.events == [] and rank._recovery_state().owner is None + assert "secondary exchange" in "\n".join(primary.__notes__) + + +def test_peer_forward_failure_prevents_release(rig): + rank, cuda, _, _ = rig + rank._recovery_reduce = lambda values, **kwargs: [1.0, 1.0] + with pytest.raises(RuntimeError, match="another rank"): + rank._release_cached_memory_for_backward(plan(True)) + assert cuda.events == [] and rank._recovery_state().owner is None + + +def test_globally_no_grad_stops_after_the_shared_status_vote(rig): + rank, cuda, _, _ = rig + calls = [] + + def reduce(values, *, op, sync_across_dp): + calls.append((op, len(values), sync_across_dp)) + return values + + rank._recovery_reduce = reduce + rank._release_cached_memory_for_backward(plan(False)) + assert calls == [("MAX", 2, True)] and cuda.events == [] + assert not rank._recovery_state().first_consumed + assert rank._recovery_state().cost > 0 + + +def test_busy_owner_is_not_overwritten_or_cleared(rig): + rank, cuda, _, _ = rig + state = rank._recovery_state() + owner = state.owner = object() + cuda.free = 1 + rank._release_cached_memory_for_backward(plan(True)) + assert state.owner is owner and state.invalid + assert "release" not in cuda.events + + +def test_new_handoff_failure_retires_only_owned_output_aliases(monkeypatch): + rank = TrainerRank(_runtime()) + refs = [] + + def forward(plan, **kwargs): + rank.device = torch.device("cuda:1") + tensor = torch.ones(2, requires_grad=True) + refs.append(weakref.ref(tensor)) + return [ForwardOutput(tensor, None, None, None)] + + _stub_forward(monkeypatch, rank, forward) + state = allocator(monkeypatch) + primary = RuntimeError("post-release sample") + state["failure"] = primary + with torch.no_grad(): + iterator = rank.forward_micro_batches([_target_request(1)], no_grad=False) + with pytest.raises(RuntimeError) as caught: + next(iterator) + assert caught.value is primary and not torch.is_grad_enabled() + assert len(refs) == 1 and refs[0]() is None + assert getattr(iterator, "gi_frame") is None + assert rank._recovery_state().owner is None diff --git a/tests/unit/test_trainer_rank_physical_reserve.py b/tests/unit/test_trainer_rank_physical_reserve.py new file mode 100644 index 000000000..bbf7ac3ef --- /dev/null +++ b/tests/unit/test_trainer_rank_physical_reserve.py @@ -0,0 +1,317 @@ +"""CPU allocator/iterator contracts; no external-library reserve calibration.""" + +from contextlib import contextmanager +from dataclasses import replace +import inspect +from typing import Any + +import pytest +import torch + +from art.trainer_rank import ForwardInput, ForwardOutput, TrainerRank, _impl +from tests.unit.test_trainer_rank_validation import ( + _runtime, + _stub_forward, + _target_request, +) + + +def allocator( + monkeypatch, + *, + free=1, + total=1000, + allocated=400, + reserved=900, + after=100, + after_reserved=None, + backend="native", +): + state: dict[str, Any] = dict( + free=free, + total=total, + allocated=allocated, + reserved=reserved, + reads=0, + releases=0, + current=torch.device("cuda:7"), + failure=None, + ) + + def info(device): + assert device == torch.device("cuda:1") + state["reads"] += 1 + if state["failure"] is not None and state["releases"]: + raise state["failure"] + return state["free"], state["total"] + + @contextmanager + def device(target): + previous, state["current"] = state["current"], target + try: + yield + finally: + state["current"] = previous + + def release(): + assert state["current"] == torch.device("cuda:1") + state["releases"] += 1 + state.update( + free=after, + reserved=allocated if after_reserved is None else after_reserved, + ) + + monkeypatch.setattr(torch.cuda, "is_available", lambda: True) + monkeypatch.setattr(torch.cuda, "get_allocator_backend", lambda: backend) + monkeypatch.setattr(torch.cuda, "mem_get_info", info) + monkeypatch.setattr( + torch.cuda, "memory_allocated", lambda _device: state["allocated"] + ) + monkeypatch.setattr( + torch.cuda, "memory_reserved", lambda _device: state["reserved"] + ) + monkeypatch.setattr(torch.cuda, "device", device) + monkeypatch.setattr(torch.cuda, "empty_cache", release) + return state + + +@pytest.mark.parametrize( + ("free", "allocated", "reserved", "after", "releases"), + [ + (29, 400, 900, 100, 1), + (30, 400, 900, 100, 0), + (31, 400, 900, 100, 0), + (1, 400, 400, 100, 0), + (1, 400, 399, 100, 0), + (1, 400, 900, 2, 1), + ], +) +def test_physical_reserve_is_only_a_soft_release_trigger( + monkeypatch, free, allocated, reserved, after, releases +): + rank = TrainerRank(_runtime()) + plan = rank._plan_flat_forward([_target_request(1)]) + rank.device = torch.device("cuda:1") + state = allocator( + monkeypatch, free=free, allocated=allocated, reserved=reserved, after=after + ) + # Native admission credits only physical free memory; handoff samples later. + check = rank._memory_check_required(1) + assert check.available_bytes == max(0, free - 30) + before = state["reads"] + rank._release_cached_memory_for_backward(plan) + assert state["releases"] == releases + assert state["reads"] - before == 1 + 2 * releases + assert state["current"] == torch.device("cuda:7") + # An unsuccessful attempt to meet the soft trigger adds no refusal/retry. + if releases: + assert state["free"] == after + + +@pytest.mark.parametrize("case", ["cpu", "all_no_grad", "inactive", "unavailable"]) +def test_non_gradient_or_non_cuda_work_does_not_query_physical_memory( + monkeypatch, case +): + rank = TrainerRank(_runtime()) + requests = ( + [ForwardInput(input_tokens=torch.tensor([1]))] + if case == "inactive" + else [_target_request(1)] + ) + with torch.set_grad_enabled(case != "all_no_grad"): + plan = rank._plan_flat_forward(requests) + rank.device = torch.device("cpu" if case == "cpu" else "cuda:1") + state = allocator(monkeypatch) + if case == "unavailable": + monkeypatch.setattr(torch.cuda, "is_available", lambda: False) + rank._release_cached_memory_for_backward(plan) + assert state["reads"] == state["releases"] == 0 + + +def test_mixed_gradient_handoff_preserves_actual_check_profile_and_context(monkeypatch): + rank = TrainerRank(_runtime()) + checks, profiles, executed, phases = [], [], [], [] + select = rank._select_next_micro_batch + + def selected(*args, **kwargs): + candidate = select(*args, **kwargs) + checks.append(candidate.check) + return candidate + + def forward(plan, **kwargs): + assert kwargs["check"] is checks[-1] + executed.append(plan) + rank.device = torch.device("cuda:1") + return [ + ForwardOutput(torch.ones(2, requires_grad=g.grad_enabled), None, None, None) + for g in plan.groups + ] + + _stub_forward(monkeypatch, rank, forward, profiled=True) + monkeypatch.setattr(rank, "_select_next_micro_batch", selected) + monkeypatch.setattr( + rank, + "_update_peak_memory_profile", + lambda plan, baseline: profiles.append((plan, baseline)), + ) + state = allocator(monkeypatch) + original_phase = _impl._telemetry_phase + + @contextmanager + def phase(name, evidence, **kwargs): + with original_phase(name, evidence, **kwargs): + yield + if name == "gradient_handoff_cache_release": + phases.append(evidence.copy()) + + monkeypatch.setattr(_impl, "_telemetry_phase", phase) + inputs = [ + [ + replace(_target_request(1), no_grad=True), + replace(_target_request(3), no_grad=False), + ] + ] + with torch.no_grad(): + iterator = rank.forward_micro_batches(inputs, no_grad=True) + batch = next(iterator) + assert not torch.is_grad_enabled() + assert [g.grad_enabled for g in executed[0].groups] == [False, True] + assert state["releases"] == 1 and state["current"] == torch.device("cuda:7") + assert batch.inputs == inputs and batch.indices == (0,) + assert ( + batch.stats.estimated_required_bytes == checks[0].estimated_required_bytes + ) + assert batch.stats.available_bytes == checks[0].available_bytes + assert ( + rank.last_forward_telemetry()["usable_limit_bytes"] + == checks[0].available_bytes + ) + target = batch.outputs[0][1].target_logprobs + target.backward(torch.ones_like(target)) + assert profiles == [] + with pytest.raises(StopIteration): + next(iterator) + assert not torch.is_grad_enabled() + assert len(checks) == len(executed) == len(profiles) == 1 + assert profiles == [(executed[0], None)] + assert phases[0]["physical_free_before_bytes"] == 1 + assert phases[0]["physical_free_after_bytes"] == 100 + assert phases[0]["reserve_trigger_bytes"] == 30 + + +@pytest.mark.parametrize("fault", ["release", "after_snapshot"]) +def test_handoff_failure_preserves_error_and_restores_device_and_grad_context( + monkeypatch, fault +): + rank = TrainerRank(_runtime()) + original = RuntimeError("original CUDA failure") + state = allocator(monkeypatch) + + def forward(plan, **kwargs): + rank.device = torch.device("cuda:1") + return [ForwardOutput(torch.ones(2, requires_grad=True), None, None, None)] + + _stub_forward(monkeypatch, rank, forward) + if fault == "release": + + def fail(): + assert state["current"] == rank.device + raise original + + monkeypatch.setattr(torch.cuda, "empty_cache", fail) + else: + state["failure"] = original + with torch.no_grad(): + iterator = rank.forward_micro_batches([_target_request(1)], no_grad=False) + with pytest.raises(RuntimeError) as caught: + next(iterator) + assert caught.value is original + assert not torch.is_grad_enabled() + assert state["current"] == torch.device("cuda:7") + assert inspect.isgenerator(iterator) + assert inspect.getgeneratorstate(iterator) == inspect.GEN_CLOSED + + +@pytest.mark.parametrize("backend", ["cudaMallocAsync", "unknown"]) +def test_unqualified_allocator_backend_skips_physical_queries(monkeypatch, backend): + rank = TrainerRank(_runtime()) + plan = rank._plan_flat_forward([_target_request(1)]) + rank.device = torch.device("cuda:1") + state = allocator(monkeypatch, backend=backend) + rank._release_cached_memory_for_backward(plan) + assert state["reads"] == state["releases"] == 0 + + +def test_matched_h200_snapshot_releases_cache_without_repricing_admission(monkeypatch): + rank = TrainerRank(_runtime()) + plan = rank._plan_flat_forward([_target_request(1)]) + rank.device = torch.device("cuda:1") + state = allocator( + monkeypatch, + free=14_352_384, + total=150_121_021_440, + allocated=92_197_523_968, + reserved=147_394_134_016, + after=42_768_990_208, + after_reserved=104_639_496_192, + ) + # Replays measured allocator counters, not a calibrated cuBLAS requirement. + check = rank._memory_check_required(7_976_316_880) + assert not check.fits # This old cached-credit admission is now refused. + original_check = (check.estimated_required_bytes, check.available_bytes, check.fits) + rank._release_cached_memory_for_backward(plan) + assert state["releases"] == 1 + assert state["free"] - 14_352_384 == 42_754_637_824 + assert 147_394_134_016 - state["reserved"] == 42_754_637_824 + assert state["allocated"] == 92_197_523_968 + assert ( + check.estimated_required_bytes, + check.available_bytes, + check.fits, + ) == original_check + + +@pytest.mark.parametrize("delta, releases", [(-1, 1), (0, 0), (1, 0)]) +def test_existing_reserve_boundary_on_h200(monkeypatch, delta, releases): + rank = TrainerRank(_runtime()) + plan = rank._plan_flat_forward([_target_request(1)]) + rank.device = torch.device("cuda:1") + total = 150_121_021_440 + reserve = int(total * _impl._MEMORY_RESERVE_FRACTION) + assert reserve == 4_503_630_643 + state = allocator(monkeypatch, total=total, free=reserve + delta) + rank._release_cached_memory_for_backward(plan) + assert state["releases"] == releases + + +def test_split_plan_releases_once_for_the_whole_gradient_handoff(monkeypatch): + rank = TrainerRank(_runtime()) + parts = tuple( + rank._plan_flat_forward([replace(_target_request(i), no_grad=no_grad)]) + for i, no_grad in [(1, True), (2, False), (3, False)] + ) + plan = _impl._SplitForwardPlan(parts, ((0,), (1,), (2,)), 3) + assert [g.grad_enabled for g in plan.groups] == [False, True, True] + rank.device = torch.device("cuda:1") + state = allocator(monkeypatch) + rank._release_cached_memory_for_backward(plan) + assert state["releases"] == 1 + + +def test_direct_forward_has_no_new_handoff_policy(monkeypatch): + rank = TrainerRank(_runtime()) + executed = [] + + def forward(plan, **kwargs): + executed.append(plan) + return [ForwardOutput(torch.ones(2, requires_grad=True), None, None, None)] + + _stub_forward(monkeypatch, rank, forward) + monkeypatch.setattr( + rank, + "_release_cached_memory_for_backward", + lambda plan: pytest.fail("direct forward is outside the iterator handoff"), + ) + output = rank.dp_rank_forward([_target_request(1)])[0] + output.target_logprobs.sum().backward() + assert len(executed) == 1 diff --git a/tests/unit/test_trainer_rank_recovery_slots_distributed.py b/tests/unit/test_trainer_rank_recovery_slots_distributed.py index 3326663cb..0456b2af1 100644 --- a/tests/unit/test_trainer_rank_recovery_slots_distributed.py +++ b/tests/unit/test_trainer_rank_recovery_slots_distributed.py @@ -11,7 +11,9 @@ HEAD = Path(__file__).resolve().parents[2] / "src" -@pytest.mark.parametrize("mode", ("fit", "both", "asymmetric")) +@pytest.mark.parametrize( + "mode", ("fit", "both", "asymmetric", "handoff", "handoff-error", "handoff-cancel") +) def test_native_checkpoint_gather_after_recovery(tmp_path, mode): selected = HEAD children = [] @@ -52,7 +54,8 @@ def test_native_checkpoint_gather_after_recovery(tmp_path, mode): ] assert all(row["error"] is None for row in rows), rows assert all(row["barrier_error"] is None for row in rows), rows - assert all(row["ensures"] == 1 for row in rows), rows + expected = 0 if mode.startswith("handoff") else 1 + assert all(row["ensures"] == expected for row in rows), rows def worker(index, mode, directory): @@ -75,6 +78,45 @@ def worker(index, mode, directory): ) rank = TrainerRank.__new__(TrainerRank) rank.device = torch.device("cpu") + if mode.startswith("handoff"): + # The second real peer has no local groups, but must join every phase. + plan = SimpleNamespace( + groups=[SimpleNamespace(grad_enabled=True)] if index == 0 else [] + ) + original = ( + KeyboardInterrupt("original forward cancellation") + if mode == "handoff-cancel" + else RuntimeError("original forward failure") + ) + error = barrier_error = None + caught = None + try: + rank._release_cached_memory_for_backward( + plan, error=original if index == 0 and mode != "handoff" else None + ) + except BaseException as exc: + caught = exc + if mode == "handoff": + valid = caught is None + else: + valid = ( + caught is original + if index == 0 + else ( + isinstance(caught, RuntimeError) and "another rank" in str(caught) + ) + ) + if not valid or rank._recovery_state().owner is not None: + error = "handoff disposition or owner differs" + try: + dist.barrier() + except BaseException as exc: + barrier_error = {"type": type(exc).__name__, "message": str(exc)} + (directory / f"rank-{index}.json").write_text( + json.dumps(dict(error=error, barrier_error=barrier_error, ensures=0)) + ) + dist.destroy_process_group() + return rank._checkpoint_mutation_lock = threading.RLock() rank._checkpoint_prefetch_lock = threading.Lock() rank._checkpoint_slots = {} diff --git a/tests/unit/test_trainer_rank_split_peak.py b/tests/unit/test_trainer_rank_split_peak.py index dbfb21d54..4ed9aff00 100644 --- a/tests/unit/test_trainer_rank_split_peak.py +++ b/tests/unit/test_trainer_rank_split_peak.py @@ -183,6 +183,11 @@ def _counter_split(monkeypatch): counters: dict[str, Any] = dict(allocated=100, peak=100, resets=[], executed=0) monkeypatch.setattr(tr, "_telemetry_phase", lambda *a, **k: nullcontext()) monkeypatch.setattr(torch.cuda, "is_available", lambda: True) + # This fixture's 10,000-byte admission budget is synthetic, not a physical + # deficit. Keep its learned-floor refusal independent of cache recovery. + monkeypatch.setattr(torch.cuda, "get_allocator_backend", lambda: "native") + monkeypatch.setattr(torch.cuda, "mem_get_info", lambda _: (1_000_000, 1_000_000)) + monkeypatch.setattr(torch.cuda, "memory_reserved", lambda _: counters["allocated"]) monkeypatch.setattr(torch.cuda, "synchronize", lambda _: None) monkeypatch.setattr(torch.cuda, "memory_allocated", lambda _: counters["allocated"]) monkeypatch.setattr(torch.cuda, "max_memory_allocated", lambda _: counters["peak"]) @@ -211,10 +216,6 @@ def execute(plan): def test_completed_iterator_preserves_caller_peak_for_next_admission(monkeypatch): rank, requests, counters = _counter_split(monkeypatch) - # This fixture's 10,000-byte admission budget is synthetic, not a physical - # deficit. Keep its learned-floor refusal independent of cache recovery. - monkeypatch.setattr(torch.cuda, "get_allocator_backend", lambda: "native") - monkeypatch.setattr(torch.cuda, "mem_get_info", lambda _: (1_000_000, 1_000_000)) releases = [] monkeypatch.setattr(torch.cuda, "empty_cache", lambda: releases.append(True)) iterator = rank.forward_micro_batches([requests], yield_empty=True) From 0a0288090ef2c2a556cf5f57c2d0d46c9ea30dd9 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 22 Sep 2026 02:45:20 +0000 Subject: [PATCH 2/9] Prototype private completed-backward and observer-cost accounting --- src/art/trainer_rank/_backward_work.py | 272 +++++++++++ src/art/trainer_rank/_impl.py | 55 ++- tests/unit/test_trainer_rank_backward_work.py | 446 ++++++++++++++++++ 3 files changed, 770 insertions(+), 3 deletions(-) create mode 100644 src/art/trainer_rank/_backward_work.py create mode 100644 tests/unit/test_trainer_rank_backward_work.py diff --git a/src/art/trainer_rank/_backward_work.py b/src/art/trainer_rank/_backward_work.py new file mode 100644 index 000000000..9e8a33f8f --- /dev/null +++ b/src/art/trainer_rank/_backward_work.py @@ -0,0 +1,272 @@ +"""Private, bounded accounting for completed ordinary-engine backward phases. + +The currency is host elapsed time, including intervening engine callbacks and +scheduling, gated by a completed CUDA tail. It is not CUDA kernel time. Observer +spans are charged conservatively, including spans also timed by recovery. +""" + +from dataclasses import dataclass +from functools import wraps +import sys +import time +from typing import Any +import weakref + +import torch + +_MAX_NS = 2**63 - 1 + + +def _measured(function): + @wraps(function) + def measured(self, *args, **kwargs): + started = None + try: + started = time.perf_counter_ns() + return function(self, *args, **kwargs) + except BaseException: + # Accounting must not replace a training error, cancellation or grad. + self.disabled = True + finally: + self._charge(started) + + return measured + + +def region(function): + """Exclude original forward/recovery work; meter only our transitions.""" + @wraps(function) + def wrapped(rank, *args, **kwargs): + work = rank._backward_work() + entered = work.enter() if work is not None else False + try: + return function(rank, *args, **kwargs) + finally: + if entered: + work.leave() + + return wrapped + + +@dataclass +class _Invocation: + started: int + blocked: bool + ended: int | None = None + tail: Any = None + + +class BackwardWork: + def __init__(self, lock, device): + self.lock = lock + self.cost_ns = 0 + self.invalid = False + started = time.perf_counter_ns() + self.device = device + self.work_ns = 0 + self.disabled = device.type != "cuda" or getattr(device, "index", None) is None + self.closed = False + self.depth = 0 + self.highest_task = -1 + # Object identity is the owner generation; callbacks hold only weak refs. + self.rows: dict[int, _Invocation] = {} + self.outputs: dict[int, tuple[Any, Any]] = {} + self._charge(started) + + def _charge(self, started): + try: + ended = time.perf_counter_ns() + with self.lock: + if ( + type(started) is not int + or type(ended) is not int + or not 0 <= started <= ended <= _MAX_NS + or self.cost_ns + ended - started > _MAX_NS + ): + self.invalid = True + else: + self.cost_ns += ended - started + except BaseException: + # Preserve the positive cost already recorded, never replace by zero. + self.invalid = True + + def _ordinary(self): + module = sys.modules.get("torch._dynamo.compiled_autograd") + return not torch.compiler.is_compiling() and ( + module is None + or not ( + getattr(module, "compiled_autograd_enabled", False) + or getattr(module, "in_compiled_autograd_region", False) + ) + ) + + @_measured + def attach(self, outputs): + with self.lock: + if self.disabled or self.closed: + return + if not self._ordinary(): + self.disabled = True + return + for key, (ref, handle) in list(self.outputs.items()): + if ref() is None: + handle.remove() + del self.outputs[key] + for output in outputs: + top_k = output.top_k + for tensor in ( + output.target_logprobs, + output.logits, + output.hidden_states, + None if top_k is None else top_k.logprobs, + ): + if tensor is None or not tensor.requires_grad: + continue + if tensor.device != self.device: + self.disabled = True + return + key = id(tensor) + if key in self.outputs and self.outputs[key][0]() is tensor: + continue + if len(self.outputs) >= 256: + self.disabled = True + return + ref = weakref.ref(self) + + def hook(grad, ref=ref): + work = ref() + if work is not None: + work._start() + # Returning None preserves the original gradient object. + + self.outputs[key] = (weakref.ref(tensor), tensor.register_hook(hook)) + + @_measured + def _start(self): + with self.lock: + if self.disabled or self.closed: + return + task = torch._C._current_graph_task_id() + if not self._ordinary() or type(task) is not int or not 0 <= task < 2**31: + self.disabled = True + return + if task in self.rows: + if self.rows[task].ended is not None: + self.disabled = True + return + if task <= self.highest_task: + self.disabled = True + return + if len(self.rows) >= 8: + # Explicit bounded retirement of an already closed interval; + # never guess that a pending/failed engine has completed. + retired = next( + (key for key, row in self.rows.items() if row.ended is not None), + None, + ) + if retired is None: + self.disabled = True + return + del self.rows[retired] + self.highest_task = task + pending = [row for row in self.rows.values() if row.ended is None] + for row in pending: + row.blocked = True + self.rows[task] = _Invocation( + time.perf_counter_ns(), bool(self.depth or pending) + ) + ref = weakref.ref(self) + + def finish(): + work = ref() + if work is not None: + work._finish(task) + + torch.autograd.Variable._execution_engine.queue_callback(finish) + + @_measured + def _finish(self, task): + ended = time.perf_counter_ns() + with self.lock: + if self.disabled or self.closed: + return + row = self.rows[task] + observed = torch._C._current_graph_task_id() + if ( + not self._ordinary() + or type(observed) is not int + or observed != task + or row.ended is not None + or type(row.started) is not int + or type(ended) is not int + or not 0 <= row.started <= ended <= _MAX_NS + ): + self.disabled = True + return + tail = torch.cuda.Event(enable_timing=False) + tail.record(torch.cuda.current_stream(self.device)) + row.ended, row.tail = ended, tail + + @_measured + def enter(self): + with self.lock: + if self.depth >= 16: + self.disabled = True + return False + self.depth += 1 + for row in self.rows.values(): + if row.ended is None: + row.blocked = True + return True + + @_measured + def leave(self): + with self.lock: + if self.depth <= 0: + self.disabled = True + else: + self.depth -= 1 + + @_measured + def harvest(self): + with self.lock: + if self.disabled or self.closed: + self.rows.clear() + return + if any(row.ended is None for row in self.rows.values()): + # A failed or unresolved engine is not a quiescent frontier. + return + addition = 0 + retired = [] + for task, row in self.rows.items(): + ready = row.tail.query() + if type(ready) is not bool: + self.disabled = True + return + if ready and not row.blocked: + addition += row.ended - row.started + if ready or row.blocked: + retired.append(task) + if self.work_ns + addition > _MAX_NS: + self.disabled = True + return + self.work_ns += addition + # Preserve completed/unready rows until a later real query or bounded + # capacity retirement. A later forward never blocks a closed interval. + for task in retired: + del self.rows[task] + + @_measured + def close(self): + with self.lock: + self.closed = True + self.rows.clear() + for _, handle in self.outputs.values(): + handle.remove() + self.outputs.clear() + + def __del__(self): + try: + self.close() + except BaseException: + pass diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 2825f3d45..b44447514 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -54,6 +54,7 @@ _local_position_pairs, estimate_prefix_tree_packed_tokens, ) +from art.trainer_rank._backward_work import BackwardWork, region as _backward_region from art.trainer_rank._planner_cost import ( COEFFICIENT_VERSION_FALLBACK, ModelGeometry, @@ -576,7 +577,8 @@ class _CacheRecoveryState: first_consumed: bool = False invalid: bool = False owner: object | None = None - lock: Any = dataclass_field(default_factory=threading.Lock) + backward: BackwardWork | None = None + lock: Any = dataclass_field(default_factory=threading.RLock) @dataclass(frozen=True) @@ -2344,6 +2346,9 @@ def _forward_micro_batches( checkpoint: AdapterSelection, yield_empty: bool, ) -> Generator[MicroBatch[ForwardInputs, ForwardOutputs], None, None]: + backward = self._backward_work() + if backward is not None: + backward.harvest() items = [_materialize(item) for item in inputs] requests = list(_flatten(items)) self._validate_replicated_top_level_count(len(items), yield_empty=yield_empty) @@ -2394,6 +2399,8 @@ def _forward_micro_batches( # Do not retain our completed graph through a new handoff traceback. del tracked_outputs, flat_outputs, outputs raise + if backward is not None: + backward.attach(tracked_outputs) stop = start + candidate.stats_global_count if stop < len(items): self._last_global_micro_batch_size = max( @@ -2432,6 +2439,8 @@ def _forward_micro_batches( subforward_count=candidate.plan.subforward_count, ), ) + if backward is not None: + backward.harvest() # The caller normally runs backward while the micro-batch is yielded. # Include that peak in future planning; forward-only profiling can # otherwise admit a later micro-batch that leaves no collective or @@ -2447,6 +2456,7 @@ def _forward_micro_batches( del tracked_outputs, flat_outputs, outputs start = stop + @_backward_region def _release_cached_memory_for_backward( self, plan: _AnyForwardPlan, *, error: BaseException | None = None ) -> None: @@ -2543,6 +2553,9 @@ def dp_rank_forward( no_grad: bool | None = None, ) -> ForwardOutputs: self._guard_forward_collective("dp_rank_forward") + backward = self._backward_work() + if backward is not None: + backward.harvest() enabled = torch.is_grad_enabled() if no_grad is None else not no_grad with torch.set_grad_enabled(enabled): self._reset_planning_telemetry() @@ -2554,6 +2567,8 @@ def dp_rank_forward( tracked_outputs = self._execute_admitted_plan( plan, check=check, context="dp_rank_forward" ) + if backward is not None: + backward.attach(tracked_outputs) return _unflatten(materialized, iter(tracked_outputs)) def _execute_admitted_plan( @@ -2569,6 +2584,7 @@ def _execute_admitted_plan( ) return outputs + @_backward_region def _execute_split_plan_with_memory_tracking( self, plan: _SplitForwardPlan, *, check: _MemoryCheck, context: str ) -> tuple[list[AnyForwardOutput], int | None, int]: @@ -4802,6 +4818,7 @@ def _group_active_request_indices( ).append(index) return tuple((slot_ref, tuple(indices)) for slot_ref, indices in groups.items()) + @_backward_region def _run_flat_plan_with_memory_tracking( self, plan: _FlatForwardPlan, @@ -5197,6 +5214,7 @@ def _admission_outcome(self, local: int) -> int: dist.all_reduce(value, op=dist.ReduceOp.MIN) return int(value.item()) + @_backward_region def _recover_admission( self, search: Callable[[], Any], @@ -5345,6 +5363,25 @@ def _recovery_state(self) -> _CacheRecoveryState: state = self._cache_recovery_state = _CacheRecoveryState() return state + def _backward_work(self) -> BackwardWork | None: + state = self._recovery_state() + started = None + try: + started = time.perf_counter_ns() + with state.lock: + if state.backward is None and not state.invalid: + state.backward = BackwardWork(state.lock, self.device) + return state.backward + except BaseException: + # Unknown accounting cost must never become free recovery budget. + state.invalid = True + return None + finally: + if state.backward is not None: + # Includes lazy setup and lock wait; constructor overlap is an + # intentional conservative charge, not an exact subtraction. + state.backward._charge(started) + def _recovery_reduce( self, values: list[float], @@ -5394,27 +5431,39 @@ def _try_cache_recovery( handoff_grad: bool = False, ) -> bool: state = self._recovery_state() + backward = self._backward_work() + if backward is not None: + backward.harvest() now = self._recovery_clock() elapsed = None if now is None or started is None else now - started with state.lock: invalid = ( state.invalid + or (backward is not None and backward.invalid) or state.owner is not owner or elapsed is None or not math.isfinite(elapsed) or elapsed < 0 ) projected = state.cost if elapsed is None else state.cost + elapsed + # B survives forward rollback. O deliberately includes observer spans + # also covered by recovery; reductions retain the original topology. + work = state.work + (0.0 if backward is None else backward.work_ns / 1e9) + accounted_cost = projected + ( + 0.0 if backward is None else backward.cost_ns / 1e9 + ) invalid |= any( not math.isfinite(value) or value < 0 for value in ( projected, state.work, + work, state.high, projected + state.high, + accounted_cost + state.high, ) ) - local_cost = [0.0, 0.0] if invalid else [projected, state.high] + local_cost = [0.0, 0.0] if invalid else [accounted_cost, state.high] # SUM costs deliberately overcharges parallel ranks; unlike MAX of # lifetime costs, it cannot miss episodes with different slow ranks. costs = self._recovery_reduce( @@ -5424,7 +5473,7 @@ def _try_cache_recovery( values = self._recovery_reduce( [ float(check.estimated_required_bytes) if check is not None else 0.0, - 0.0 if invalid else state.work, + 0.0 if invalid else work, float(state.first_consumed), float(invalid), ], diff --git a/tests/unit/test_trainer_rank_backward_work.py b/tests/unit/test_trainer_rank_backward_work.py new file mode 100644 index 000000000..ba4ff6c47 --- /dev/null +++ b/tests/unit/test_trainer_rank_backward_work.py @@ -0,0 +1,446 @@ +"""Prospective bounded accounting controls; no Torch/model/device import. + +Load the actual helper with a scalar Torch facade, and extract the actual +recovery/rollback methods rather than copy their decision arithmetic. +""" + +import ast +from dataclasses import dataclass, field +import gc +import importlib.util +import math +import os +from pathlib import Path +import sys +import threading +from types import SimpleNamespace as NS +import unittest +from unittest.mock import patch +import weakref + + +ROOT = Path(__file__).resolve().parents[2] + + +class Clock: + value = 1000 + + def perf_counter_ns(self): + self.value += 1 + return self.value + + +class Tensor: + def __init__(self, device): + self.device, self.requires_grad = device, True + self.hooks = {} + + def register_hook(self, callback): + key = len(self.hooks) + self.hooks[key] = callback + ref = weakref.ref(self) + + def remove(): + tensor = ref() + if tensor is not None: + tensor.hooks.pop(key, None) + + return NS(remove=remove) + + +class CUDA: + def __init__(self): + self.events = [] + self.free, self.releases = 0, 0 + self.failure = None + + def Event(self, *, enable_timing): + assert enable_timing is False + if self.failure is not None: + raise self.failure + event = NS(ready=False, stream=None, queries=0) + + def query(): + event.queries += 1 + return event.ready + + event.query = query + event.record = lambda stream: setattr(event, "stream", stream) + self.events.append(event) + return event + + def current_stream(self, device): + return ("caller", device.index) + + def is_available(self): + return True + + def get_allocator_backend(self): + return "native" + + def mem_get_info(self, device): + return self.free, 1000 + + def memory_allocated(self, device): + return 0 + + def memory_reserved(self, device): + return 100 + + def empty_cache(self): + self.releases += 1 + self.free = 500 + + +class TestBackwardWork(unittest.TestCase): + def setUp(self): + self.clock, self.cuda = Clock(), CUDA() + self.task, self.callbacks = 0, [] + self.device = NS(type="cuda", index=0) + self.torch = NS( + cuda=self.cuda, + compiler=NS(is_compiling=lambda: False), + _C=NS(_current_graph_task_id=lambda: self.task), + autograd=NS(Variable=NS(_execution_engine=NS( + queue_callback=self.callbacks.append, + ))), + ) + name = "_art_backward_work_control" + spec = importlib.util.spec_from_file_location( + name, ROOT / "src/art/trainer_rank/_backward_work.py" + ) + self.module = importlib.util.module_from_spec(spec) + with patch.dict(sys.modules, {"torch": self.torch, name: self.module}): + spec.loader.exec_module(self.module) + self.module.time = self.clock + self.work = self.module.BackwardWork(threading.RLock(), self.device) + self.addCleanup(self.work.close) + + def output(self, tensor): + return NS(target_logprobs=tensor, logits=tensor, hidden_states=tensor, + top_k=NS(logprobs=tensor)) + + def start(self, task): + self.task = task + self.work._start() + + def finish(self, task, *, ready=True): + self.task = task + callback = self.callbacks.pop(0) + callback() + if self.cuda.events: + self.cuda.events[-1].ready = ready + + def completed(self, task=1, *, ready=True): + self.start(task) + self.clock.value += 1000 + self.finish(task, ready=ready) + return self.work.rows[task].ended - self.work.rows[task].started + + def actual_rank(self): + source = ast.parse((ROOT / "src/art/trainer_rank/_impl.py").read_text()) + names = {"_CacheRecoveryState", "_MemoryCheck"} + methods = {"_recovery_state", "_backward_work", "_try_cache_recovery", + "_execute_split_plan_with_memory_tracking"} + selected = [node for node in source.body if isinstance(node, ast.ClassDef) + and node.name in names] + trainer = next(node for node in source.body if isinstance(node, ast.ClassDef) + and node.name == "TrainerRank") + trainer.body = [node for node in trainer.body + if isinstance(node, ast.FunctionDef) and node.name in methods] + ns = dict(__name__=__name__, dataclass=dataclass, dataclass_field=field, + threading=threading, BackwardWork=self.module.BackwardWork, + _backward_region=self.module.region, torch=self.torch, + math=math, os=os, time=self.clock, cast=lambda typ, value: value, + TrainerRankMemoryError=type("MemoryRefusal", (RuntimeError,), {}), + _TEST_HOOKS_ENV="ART_BACKWARD_CONTROL_ONLY", + _TEST_MEMORY_LIMIT_ENV="ART_BACKWARD_CONTROL_LIMIT", + _MEMORY_RESERVE_FRACTION=0.05) + tree = ast.Module(body=[ast.ImportFrom(module="__future__", names=[ + ast.alias(name="annotations")], level=0), *selected, trainer], type_ignores=[]) + exec(compile(ast.fix_missing_locations(tree), "actual-accounting-methods", "exec"), ns) + rank = object.__new__(ns["TrainerRank"]) + rank.device = self.device + state = rank._recovery_state() + self.work.lock = state.lock + state.backward = self.work + rank._recovery_clock = lambda: 1.0 + rank._available_memory_bytes = lambda: self.cuda.free + return rank, state, ns + + def test_fixed_endpoint_unready_survives_later_forward_and_idle(self): + duration = self.completed(ready=False) + ended = self.work.rows[1].ended + self.work.harvest() + self.assertEqual(self.work.work_ns, 0) + self.work.enter() + self.clock.value += 10**12 + self.work.leave() + self.assertFalse(self.work.rows[1].blocked) + self.assertEqual(self.work.rows[1].ended, ended) + self.cuda.events[0].ready = True + self.work.harvest() + self.assertEqual(self.work.work_ns, duration) + self.work.harvest() + self.assertEqual(self.work.work_ns, duration) + self.assertEqual(self.cuda.events[0].stream, ("caller", 0)) + + def test_four_fields_repeated_attach_and_multiple_outputs_queue_once(self): + first, second = Tensor(self.device), Tensor(self.device) + for _ in range(2): + self.work.attach([self.output(first), self.output(second)]) + self.assertEqual((len(first.hooks), len(second.hooks)), (1, 1)) + self.task = 1 + gradient = object() + for tensor in (first, second): + self.assertIsNone(next(iter(tensor.hooks.values()))(gradient)) + self.assertEqual(len(self.callbacks), 1) + self.finish(1) + self.work.harvest() + self.assertGreater(self.work.work_ns, 0) + + def test_retained_graph_new_ids_and_retired_id_cannot_revive_credit(self): + expected = self.completed(1) + self.work.harvest() + expected += self.completed(2) + self.work.harvest() + self.assertEqual(self.work.work_ns, expected) + self.start(1) + self.assertTrue(self.work.disabled) + self.work.harvest() + self.assertEqual(self.work.work_ns, expected) + + def test_nested_and_forward_overlap_retire_zero(self): + self.start(1) + self.start(2) + # The engine callbacks may complete inner first. + self.callbacks.reverse() + self.finish(2) + self.finish(1) + self.work.harvest() + self.assertEqual(self.work.work_ns, 0) + self.work.enter() + self.completed(3) + self.work.leave() + self.work.harvest() + self.assertEqual(self.work.work_ns, 0) + self.start(4) + self.work.enter() + self.work.leave() + self.finish(4) + self.work.harvest() + self.assertEqual(self.work.work_ns, 0) + + def test_failed_pending_blocks_commit_without_clock_timeout(self): + self.start(1) # Node failure means its completion callback never arrives. + self.callbacks.clear() + self.completed(2) + self.clock.value += 10**12 + self.work.harvest() + self.assertEqual(self.work.work_ns, 0) + self.assertIn(1, self.work.rows) + + def test_bounded_unready_retirement_and_long_stream_keep_committed_total(self): + committed = self.completed(1) + self.work.harvest() + for task in range(2, 12): + self.completed(task, ready=False) + self.assertEqual(len(self.work.rows), 8) + self.assertNotIn(2, self.work.rows) + self.assertFalse(self.work.disabled) + for event in self.cuda.events: + event.ready = True + committed += sum(row.ended - row.started for row in self.work.rows.values()) + self.work.harvest() + self.assertEqual(self.work.work_ns, committed) + for task in range(12, 1012): + committed += self.completed(task) + self.work.harvest() + self.assertEqual(len(self.work.rows), 0) + self.assertEqual(self.work.work_ns, committed) + + def test_pending_overflow_preserves_prior_credit_and_positive_cost(self): + committed = self.completed(1) + self.work.harvest() + for task in range(2, 11): + self.start(task) + self.assertTrue(self.work.disabled) + self.assertLessEqual(len(self.work.rows), 8) + cost = self.work.cost_ns + self.work.harvest() + self.assertEqual(self.work.work_ns, committed) + self.assertGreaterEqual(self.work.cost_ns, cost) + + def test_original_exception_identity_cause_and_observer_fault_neutrality(self): + primary, cause = KeyboardInterrupt("original"), ValueError("cause") + rank = NS(_backward_work=lambda: self.work) + + @self.module.region + def fail(rank): + self.cuda.failure = RuntimeError("diagnostic tail") + self.start(1) + self.finish(1) + raise primary from cause + + with self.assertRaises(KeyboardInterrupt) as caught: + fail(rank) + self.assertIs(caught.exception, primary) + self.assertIs(primary.__cause__, cause) + self.assertTrue(self.work.disabled) + self.assertGreater(self.work.cost_ns, 0) + + def test_queue_failure_does_not_replace_gradient_or_erase_previous_credit(self): + expected = self.completed(1) + self.work.harvest() + tensor = Tensor(self.device) + self.work.attach([self.output(tensor)]) + self.task = 2 + with patch.object(self.torch.autograd.Variable._execution_engine, + "queue_callback", side_effect=KeyboardInterrupt("queue")): + self.assertIsNone(next(iter(tensor.hooks.values()))(object())) + self.assertTrue(self.work.disabled) + self.work.harvest() + self.assertEqual(self.work.work_ns, expected) + + def test_clock_fault_cost_overflow_and_unknown_readiness_fail_closed(self): + prior = self.work.cost_ns + with patch.object(self.clock, "perf_counter_ns", side_effect=RuntimeError("clock")): + self.work.harvest() + self.assertTrue(self.work.invalid) + self.assertEqual(self.work.cost_ns, prior) + self.work.invalid = False + self.work.cost_ns = self.module._MAX_NS + self.work.harvest() + self.assertTrue(self.work.invalid) + self.assertEqual(self.work.cost_ns, self.module._MAX_NS) + + def test_wrong_device_mode_id_and_nonbool_readiness_withhold(self): + for task in (-1, 2**31, True): + with self.subTest(task=task): + self.work.disabled = False + self.start(task) + self.assertTrue(self.work.disabled) + self.work.disabled = False + tensor = Tensor(NS(type="cuda", index=1)) + self.work.attach([self.output(tensor)]) + self.assertTrue(self.work.disabled) + self.assertFalse(tensor.hooks) + self.work.disabled = False + with patch.dict(sys.modules, {"torch._dynamo.compiled_autograd": NS( + compiled_autograd_enabled=True)}): + self.start(1) + self.assertTrue(self.work.disabled) + self.work.disabled = False + self.completed(2) + self.cuda.events[-1].ready = 1 + self.work.harvest() + self.assertTrue(self.work.disabled) + self.assertEqual(self.work.work_ns, 0) + + def test_weak_outputs_owner_and_close_remove_only_owned_hooks(self): + tensor = Tensor(self.device) + tensor.register_hook(lambda grad: grad) + self.work.attach([self.output(tensor)]) + self.assertEqual(len(tensor.hooks), 2) + self.work.close() + self.assertEqual(len(tensor.hooks), 1) + ref = weakref.ref(tensor) + del tensor + gc.collect() + self.assertIsNone(ref()) + owner = self.module.BackwardWork(threading.RLock(), self.device) + tensor = Tensor(self.device) + owner.attach([self.output(tensor)]) + ref = weakref.ref(owner) + del owner + gc.collect() + self.assertIsNone(ref()) + self.assertEqual(tensor.hooks, {}) + + def test_actual_reducer_uses_sum_cost_and_max_local_sum_not_sum_of_maxima(self): + rank, state, ns = self.actual_rank() + state.work, state.cost, state.high = 100.0, 4.0, 1.0 + state.first_consumed, state.owner = True, object() + calls = [] + + def reduce(values, *, op, sync_across_dp): + calls.append((op, list(values))) + if op == "SUM": + return [values[0] + 3.0, values[1] + 1.0] + if op == "MAX": + # Other rank has F=0, B=100; MAX(local F+B)=100, not 200. + return [values[0], max(values[1], 100.0), values[2], values[3]] + return values + + rank._recovery_reduce = reduce + result = rank._try_cache_recovery(ns["_MemoryCheck"](80, 0, False), + sync_across_dp=True, owner=state.owner, started=0.0) + self.assertFalse(result) + self.assertEqual([op for op, _ in calls], ["SUM", "MAX", "MIN"]) + self.assertEqual(calls[1][1][1], 100.0) + self.assertGreater(calls[0][1][0], 5.0) # Original C+elapsed plus measured O. + self.assertEqual(self.cuda.releases, 0) + + def test_actual_first_repeat_rule_and_invalid_meter(self): + rank, state, ns = self.actual_rank() + rank._recovery_reduce = lambda values, **kw: values + state.owner, state.high = object(), 0.1 + check = ns["_MemoryCheck"](80, 0, False) + self.assertTrue(rank._try_cache_recovery(check, sync_across_dp=False, + owner=state.owner, started=0.0)) # Original first free release. + self.cuda.free = 0 + self.assertFalse(rank._try_cache_recovery(check, sync_across_dp=False, + owner=state.owner, started=0.0)) + # Use a genuinely completed host interval in this scalar fixture. + self.start(1) + self.clock.value += 30_000_000_000 + self.finish(1) + self.assertTrue(rank._try_cache_recovery(check, sync_across_dp=False, + owner=state.owner, started=0.0)) + self.cuda.free = 0 + self.work.invalid = True + self.assertFalse(rank._try_cache_recovery(check, sync_across_dp=False, + owner=state.owner, started=0.0)) + + def test_actual_split_rollback_keeps_b_and_o_and_primary_error(self): + rank, state, ns = self.actual_rank() + self.completed() + self.work.harvest() + credited, cost = self.work.work_ns, self.work.cost_ns + state.work = 5.0 + primary = ValueError("second subforward") + count = 0 + + def execute(*args, **kw): + nonlocal count + count += 1 + state.work += 10.0 + if count == 2: + raise primary + return [object()], None + + rank._run_flat_plan_with_memory_tracking = execute + plan = NS(request_count=2, subforwards=[object(), object()], + request_indices=[[0], [1]], subforward_count=2) + with self.assertRaises(ValueError) as caught: + rank._execute_split_plan_with_memory_tracking(plan, check=None, context="test") + self.assertIs(caught.exception, primary) + self.assertEqual(state.work, 5.0) + self.assertEqual(self.work.work_ns, credited) + self.assertGreater(self.work.cost_ns, cost) + + def test_original_invalid_forward_guard_cannot_be_masked_by_positive_b(self): + rank, state, ns = self.actual_rank() + rank._recovery_reduce = lambda values, **kw: values + state.owner, state.work, state.high = object(), -1.0, 0.1 + self.start(1) + self.clock.value += 30_000_000_000 + self.finish(1) + self.assertFalse(rank._try_cache_recovery(ns["_MemoryCheck"](80, 0, False), + sync_across_dp=False, owner=state.owner, started=0.0)) + self.assertTrue(state.invalid) + self.assertEqual(self.cuda.releases, 0) + + +if __name__ == "__main__": + unittest.main() From 0ea9f8cee7f8d234ae8a477a74dd60d5adae23f5 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 22 Sep 2026 02:55:59 +0000 Subject: [PATCH 3/9] Preserve cancellation identity in prospective backward accounting --- src/art/trainer_rank/_backward_work.py | 36 +++-- src/art/trainer_rank/_impl.py | 9 +- tests/unit/test_trainer_rank_backward_work.py | 123 +++++++++++++++++- 3 files changed, 151 insertions(+), 17 deletions(-) diff --git a/src/art/trainer_rank/_backward_work.py b/src/art/trainer_rank/_backward_work.py index 9e8a33f8f..647745de8 100644 --- a/src/art/trainer_rank/_backward_work.py +++ b/src/art/trainer_rank/_backward_work.py @@ -21,14 +21,18 @@ def _measured(function): @wraps(function) def measured(self, *args, **kwargs): started = None + primary = False try: started = time.perf_counter_ns() return function(self, *args, **kwargs) + except Exception: + self.disabled = True except BaseException: - # Accounting must not replace a training error, cancellation or grad. self.disabled = True + primary = True + raise finally: - self._charge(started) + self._charge(started, primary=primary) return measured @@ -39,11 +43,19 @@ def region(function): def wrapped(rank, *args, **kwargs): work = rank._backward_work() entered = work.enter() if work is not None else False + primary = False try: return function(rank, *args, **kwargs) + except BaseException: + primary = True + raise finally: - if entered: - work.leave() + try: + if entered: + work.leave() + except BaseException: + if not primary: + raise return wrapped @@ -73,7 +85,7 @@ def __init__(self, lock, device): self.outputs: dict[int, tuple[Any, Any]] = {} self._charge(started) - def _charge(self, started): + def _charge(self, started, *, primary=False): try: ended = time.perf_counter_ns() with self.lock: @@ -86,9 +98,15 @@ def _charge(self, started): self.invalid = True else: self.cost_ns += ended - started - except BaseException: + except Exception: # Preserve the positive cost already recorded, never replace by zero. self.invalid = True + except BaseException: + self.invalid = True + # A fresh cancellation propagates. Only a secondary meter failure + # may be suppressed while preserving an already propagating error. + if not primary: + raise def _ordinary(self): module = sys.modules.get("torch._dynamo.compiled_autograd") @@ -264,9 +282,3 @@ def close(self): for _, handle in self.outputs.values(): handle.remove() self.outputs.clear() - - def __del__(self): - try: - self.close() - except BaseException: - pass diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index b44447514..7ff117806 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -5366,21 +5366,26 @@ def _recovery_state(self) -> _CacheRecoveryState: def _backward_work(self) -> BackwardWork | None: state = self._recovery_state() started = None + primary = False try: started = time.perf_counter_ns() with state.lock: if state.backward is None and not state.invalid: state.backward = BackwardWork(state.lock, self.device) return state.backward - except BaseException: + except Exception: # Unknown accounting cost must never become free recovery budget. state.invalid = True return None + except BaseException: + state.invalid = True + primary = True + raise finally: if state.backward is not None: # Includes lazy setup and lock wait; constructor overlap is an # intentional conservative charge, not an exact subtraction. - state.backward._charge(started) + state.backward._charge(started, primary=primary) def _recovery_reduce( self, diff --git a/tests/unit/test_trainer_rank_backward_work.py b/tests/unit/test_trainer_rank_backward_work.py index ba4ff6c47..c0f10fe37 100644 --- a/tests/unit/test_trainer_rank_backward_work.py +++ b/tests/unit/test_trainer_rank_backward_work.py @@ -5,6 +5,7 @@ """ import ast +from asyncio import CancelledError from dataclasses import dataclass, field import gc import importlib.util @@ -289,19 +290,131 @@ def fail(rank): self.assertTrue(self.work.disabled) self.assertGreater(self.work.cost_ns, 0) - def test_queue_failure_does_not_replace_gradient_or_erase_previous_credit(self): + def test_ordinary_queue_fault_does_not_replace_gradient_or_erase_previous_credit(self): expected = self.completed(1) self.work.harvest() tensor = Tensor(self.device) self.work.attach([self.output(tensor)]) self.task = 2 with patch.object(self.torch.autograd.Variable._execution_engine, - "queue_callback", side_effect=KeyboardInterrupt("queue")): + "queue_callback", side_effect=RuntimeError("queue")): self.assertIsNone(next(iter(tensor.hooks.values()))(object())) self.assertTrue(self.work.disabled) self.work.harvest() self.assertEqual(self.work.work_ns, expected) + def test_new_cancellation_at_each_observer_boundary_propagates_exact_object(self): + original_work = self.work + for error_type in (KeyboardInterrupt, SystemExit, GeneratorExit, CancelledError): + for stage in ("attach", "hook", "callback", "harvest", "leave", "close"): + with self.subTest(error=error_type, stage=stage): + self.work = self.module.BackwardWork(threading.RLock(), self.device) + self.callbacks.clear() + self.task = 1 + primary, cause = error_type(stage), ValueError("original cause") + primary.__cause__ = cause + tensor = Tensor(self.device) + if stage == "attach": + target, name = tensor, "register_hook" + action = lambda: self.work.attach([self.output(tensor)]) + elif stage == "hook": + self.work.attach([self.output(tensor)]) + target = self.torch.autograd.Variable._execution_engine + name = "queue_callback" + action = lambda: next(iter(tensor.hooks.values()))(object()) + elif stage == "callback": + self.start(1) + target, name = self.cuda, "Event" + action = self.callbacks.pop(0) + elif stage == "harvest": + self.completed(1) + target, name = self.cuda.events[-1], "query" + action = self.work.harvest + elif stage == "leave": + self.work.enter() + target, name = self.clock, "perf_counter_ns" + action = self.work.leave + else: + self.work.attach([self.output(tensor)]) + target, name = next(iter(self.work.outputs.values()))[1], "remove" + action = self.work.close + with patch.object(target, name, side_effect=primary): + with self.assertRaises(BaseException) as caught: + action() + self.assertIs(caught.exception, primary) + self.assertIs(primary.__cause__, cause) + self.assertTrue(self.work.disabled) + self.assertEqual(self.work.work_ns, 0) + self.work.close() + self.work = original_work + + def test_meter_propagates_new_cancellation_and_preserves_inflight_primary(self): + primary, secondary = KeyboardInterrupt("queue"), SystemExit("meter") + cause = ValueError("cause") + primary.__cause__ = cause + self.task = 1 + with patch.object(self.clock, "perf_counter_ns", side_effect=[2000, 2001, secondary]): + with patch.object(self.torch.autograd.Variable._execution_engine, + "queue_callback", side_effect=primary): + with self.assertRaises(KeyboardInterrupt) as caught: + self.work._start() + self.assertIs(caught.exception, primary) + self.assertIs(primary.__cause__, cause) + self.assertTrue(self.work.invalid) + # With no primary, a newly delivered meter cancellation must escape. + for action in (lambda: self.work._charge(2000), self.work.harvest): + self.work.disabled = False + self.work.rows.clear() + clock = [secondary] if action != self.work.harvest else [3000, secondary] + with patch.object(self.clock, "perf_counter_ns", side_effect=clock): + with self.assertRaises(SystemExit) as caught: + action() + self.assertIs(caught.exception, secondary) + + def test_actual_lookup_and_constructor_cancellation_identity(self): + rank, state, ns = self.actual_rank() + primary, secondary = CancelledError("lookup"), GeneratorExit("meter") + cause = ValueError("lookup cause") + primary.__cause__ = cause + with patch.object(self.clock, "perf_counter_ns", side_effect=[primary, secondary]): + with self.assertRaises(CancelledError) as caught: + rank._backward_work() + self.assertIs(caught.exception, primary) + self.assertIs(primary.__cause__, cause) + self.assertTrue(state.invalid) + self.assertTrue(self.work.invalid) + state.invalid, state.backward = False, None + + def construction(*args): + raise primary + + with patch.dict(ns, BackwardWork=construction): + with self.assertRaises(CancelledError) as caught: + rank._backward_work() + self.assertIs(caught.exception, primary) + self.assertIs(primary.__cause__, cause) + self.assertTrue(state.invalid) + + def test_region_cleanup_preserves_body_error_but_new_cancellation_escapes(self): + rank = NS(_backward_work=lambda: self.work) + secondary = KeyboardInterrupt("leave") + for primary in (None, ValueError("body"), CancelledError("body cancellation")): + with self.subTest(primary=primary): + cause = RuntimeError("body cause") + + @self.module.region + def body(rank): + if primary is not None: + raise primary from cause + return "original result" + + with patch.object(self.work, "leave", side_effect=secondary): + with self.assertRaises(BaseException) as caught: + body(rank) + self.assertIs(caught.exception, secondary if primary is None else primary) + if primary is not None: + self.assertIs(primary.__cause__, cause) + def test_clock_fault_cost_overflow_and_unknown_readiness_fail_closed(self): prior = self.work.cost_ns with patch.object(self.clock, "perf_counter_ns", side_effect=RuntimeError("clock")): @@ -355,7 +468,11 @@ def test_weak_outputs_owner_and_close_remove_only_owned_hooks(self): del owner gc.collect() self.assertIsNone(ref()) - self.assertEqual(tensor.hooks, {}) + self.assertNotIn("__del__", vars(self.module.BackwardWork)) + # No Python finalizer can intercept cancellation. Caller-retained tensors + # may keep a bounded weak hook; it is inert after the rank owner is gone. + self.assertEqual(len(tensor.hooks), 1) + self.assertIsNone(next(iter(tensor.hooks.values()))(object())) def test_actual_reducer_uses_sum_cost_and_max_local_sum_not_sum_of_maxima(self): rank, state, ns = self.actual_rank() From d129f7249ee8ccf9a15b6d3c0ec180aaf4a3682c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 22 Sep 2026 03:07:48 +0000 Subject: [PATCH 4/9] Make backward exclusion transition cleanup cancellation-safe --- src/art/trainer_rank/_backward_work.py | 78 ++++++++++++------- tests/unit/test_trainer_rank_backward_work.py | 55 ++++++++++--- 2 files changed, 92 insertions(+), 41 deletions(-) diff --git a/src/art/trainer_rank/_backward_work.py b/src/art/trainer_rank/_backward_work.py index 647745de8..9715b6982 100644 --- a/src/art/trainer_rank/_backward_work.py +++ b/src/art/trainer_rank/_backward_work.py @@ -6,6 +6,7 @@ """ from dataclasses import dataclass +from contextlib import contextmanager, nullcontext from functools import wraps import sys import time @@ -42,20 +43,8 @@ def region(function): @wraps(function) def wrapped(rank, *args, **kwargs): work = rank._backward_work() - entered = work.enter() if work is not None else False - primary = False - try: + with work.region() if work is not None else nullcontext(): return function(rank, *args, **kwargs) - except BaseException: - primary = True - raise - finally: - try: - if entered: - work.leave() - except BaseException: - if not primary: - raise return wrapped @@ -225,25 +214,54 @@ def _finish(self, task): tail.record(torch.cuda.current_stream(self.device)) row.ended, row.tail = ended, tail - @_measured - def enter(self): - with self.lock: - if self.depth >= 16: + @contextmanager + def region(self): + """Own one exclusion count, even when a transition clock cancels.""" + started, entered, primary = None, False, False + try: + try: + started = time.perf_counter_ns() + with self.lock: + if self.depth >= 16: + self.disabled = True + else: + self.depth += 1 + entered = True + for row in self.rows.values(): + if row.ended is None: + row.blocked = True + except Exception: self.disabled = True - return False - self.depth += 1 - for row in self.rows.values(): - if row.ended is None: - row.blocked = True - return True - - @_measured - def leave(self): - with self.lock: - if self.depth <= 0: + except BaseException: self.disabled = True - else: - self.depth -= 1 + primary = True + raise + finally: + self._charge(started, primary=primary) + yield + except BaseException: + primary = True + raise + finally: + if entered: + started = None + try: + try: + started = time.perf_counter_ns() + finally: + # The exit clock must not prevent owned cleanup. Never + # restore a saved depth over another thread's region. + with self.lock: + self.depth -= 1 + except Exception: + self.disabled = True + except BaseException: + self.disabled = True + if not primary: + primary = True + raise + finally: + self._charge(started, primary=primary) @_measured def harvest(self): diff --git a/tests/unit/test_trainer_rank_backward_work.py b/tests/unit/test_trainer_rank_backward_work.py index c0f10fe37..75f7646b5 100644 --- a/tests/unit/test_trainer_rank_backward_work.py +++ b/tests/unit/test_trainer_rank_backward_work.py @@ -174,9 +174,11 @@ def test_fixed_endpoint_unready_survives_later_forward_and_idle(self): ended = self.work.rows[1].ended self.work.harvest() self.assertEqual(self.work.work_ns, 0) - self.work.enter() - self.clock.value += 10**12 - self.work.leave() + cost = self.work.cost_ns + with self.work.region(): + self.clock.value += 10**12 + self.assertLess(self.work.cost_ns - cost, 1000) + self.assertEqual(self.work.depth, 0) self.assertFalse(self.work.rows[1].blocked) self.assertEqual(self.work.rows[1].ended, ended) self.cuda.events[0].ready = True @@ -220,14 +222,13 @@ def test_nested_and_forward_overlap_retire_zero(self): self.finish(1) self.work.harvest() self.assertEqual(self.work.work_ns, 0) - self.work.enter() - self.completed(3) - self.work.leave() + with self.work.region(): + self.completed(3) self.work.harvest() self.assertEqual(self.work.work_ns, 0) self.start(4) - self.work.enter() - self.work.leave() + with self.work.region(): + pass self.finish(4) self.work.harvest() self.assertEqual(self.work.work_ns, 0) @@ -289,6 +290,7 @@ def fail(rank): self.assertIs(primary.__cause__, cause) self.assertTrue(self.work.disabled) self.assertGreater(self.work.cost_ns, 0) + self.assertEqual(self.work.depth, 0) def test_ordinary_queue_fault_does_not_replace_gradient_or_erase_previous_credit(self): expected = self.completed(1) @@ -331,9 +333,10 @@ def test_new_cancellation_at_each_observer_boundary_propagates_exact_object(self target, name = self.cuda.events[-1], "query" action = self.work.harvest elif stage == "leave": - self.work.enter() + scope = self.work.region() + scope.__enter__() target, name = self.clock, "perf_counter_ns" - action = self.work.leave + action = lambda: scope.__exit__(None, None, None) else: self.work.attach([self.output(tensor)]) target, name = next(iter(self.work.outputs.values()))[1], "remove" @@ -345,6 +348,7 @@ def test_new_cancellation_at_each_observer_boundary_propagates_exact_object(self self.assertIs(primary.__cause__, cause) self.assertTrue(self.work.disabled) self.assertEqual(self.work.work_ns, 0) + self.assertEqual(self.work.depth, 0) self.work.close() self.work = original_work @@ -408,12 +412,41 @@ def body(rank): raise primary from cause return "original result" - with patch.object(self.work, "leave", side_effect=secondary): + with patch.object(self.clock, "perf_counter_ns", + side_effect=[10000, 10001, secondary, 10003]): with self.assertRaises(BaseException) as caught: body(rank) self.assertIs(caught.exception, secondary if primary is None else primary) + self.assertEqual(self.work.depth, 0) if primary is not None: self.assertIs(primary.__cause__, cause) + # Enter has incremented its own count when its final charge cancels. + # Cleanup must remove only that count, retaining an existing outer one. + called = [] + + @self.module.region + def untouched(rank): + called.append(True) + + with self.work.region(): + self.assertEqual(self.work.depth, 1) + with patch.object(self.clock, "perf_counter_ns", + side_effect=[20000, secondary, 20002, 20003]): + with self.assertRaises(KeyboardInterrupt) as caught: + untouched(rank) + self.assertIs(caught.exception, secondary) + self.assertEqual(self.work.depth, 1) + self.assertEqual(called, []) + scope = self.work.region() + scope.__enter__() + self.assertEqual(self.work.depth, 2) + another = SystemExit("secondary exit meter") + with patch.object(self.clock, "perf_counter_ns", side_effect=[secondary, another]): + with self.assertRaises(KeyboardInterrupt) as caught: + scope.__exit__(None, None, None) + self.assertIs(caught.exception, secondary) + self.assertEqual(self.work.depth, 1) + self.assertEqual(self.work.depth, 0) def test_clock_fault_cost_overflow_and_unknown_readiness_fail_closed(self): prior = self.work.cost_ns From 2fee9bea8aa2d15feb91e98e603f354a0886345b Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 22 Sep 2026 05:31:51 +0000 Subject: [PATCH 5/9] Test backward accounting across checkpoint recomputation --- .../test_dispatcher_graph_retention.py | 249 ++++++++++++++++++ 1 file changed, 249 insertions(+) diff --git a/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py b/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py index fbf273892..3968367a0 100644 --- a/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py +++ b/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py @@ -4,6 +4,7 @@ from functools import partial import gc import pickle +import threading from types import SimpleNamespace from typing import Any, cast import weakref @@ -11,6 +12,7 @@ import pytest import torch from torch._dynamo.testing import CompileCounterWithBackend +import torch.utils.checkpoint as torch_checkpoint pytest.importorskip("megatron.bridge") @@ -21,7 +23,9 @@ MoEFlexTokenDispatcher, ) +from art.megatron import lora as lora_module from art.trainer_rank import TrainerRank +from art.trainer_rank._backward_work import BackwardWork from art.trainer_rank._impl import _configure_moe_dispatcher_caches @@ -125,6 +129,251 @@ def cpu_checkpoint_rng(monkeypatch): ) +@pytest.fixture +def cpu_backward_events(cpu_checkpoint_rng, monkeypatch): + # Synthetic readiness tests engine bookkeeping, not CUDA completion. Keep + # production's CPU-disabled default unless a case explicitly enables it. + assert not torch.cuda.is_initialized() + assert getattr(mcore_random.checkpoint, "_art_lora_slot_context_patch", False) + assert getattr(torch_checkpoint.checkpoint, "_art_lora_slot_context_patch", False) + events = [] + stream = object() + + class ReadyEvent: + def __init__(self, *, enable_timing): + assert not enable_timing + events.append(self) + + def record(self, actual_stream): + assert actual_stream is stream + self.task = torch._C._current_graph_task_id() + + def query(self): + return True + + monkeypatch.setattr(torch.cuda, "Event", ReadyEvent) + monkeypatch.setattr(torch.cuda, "current_stream", lambda device: stream) + # MCore leaves this flag set when recomputation raises. Restore the caller's + # original state at fixture teardown, without changing that failure path. + monkeypatch.setattr(mcore_random, "IS_CHECKPOINTING", False) + yield events + assert not torch.cuda.is_initialized() + + +def _backward_inputs(): + x = ( + torch.linspace(-0.5, 0.5, 24, dtype=torch.float64) + .reshape(4, 6) + .requires_grad_() + ) + params = [ + torch.linspace(-0.2 + i * 0.03, 0.2 + i * 0.03, 12, dtype=torch.float64) + .reshape(6, 2) + .requires_grad_() + for i in range(4) + ] + return x, params + + +def _backward_layer(value, a, b): + return torch.tanh(value + 0.25 * ((value @ a) @ b.T)) + + +def _backward_outputs(hidden, head): + output = SimpleNamespace( + target_logprobs=head[:, 0], + logits=head[:, :3], + hidden_states=hidden, + top_k=SimpleNamespace(logprobs=head[:, 3:]), + ) + fields = (output.target_logprobs, output.logits, hidden, output.top_k.logprobs) + return output, sum(t.square().mean() + 0.07 * t.sum() for t in fields) + + +@pytest.mark.parametrize( + "enabled", [False, True], ids=["cpu-disabled", "cpu-ready-facade"] +) +def test_backward_work_mixed_checkpoint(cpu_backward_events, enabled): + work = BackwardWork(threading.RLock(), torch.device("cpu")) + assert work.disabled + if enabled: + work.disabled = False + owned = lora_module.LoRASlotRef("lora", "owned") + ambient = lora_module.LoRASlotRef("lora", "ambient") + + def iteration(): + x, params = _backward_inputs() + reference_x, reference_params = _backward_inputs() + calls, inner_tasks, recomputed = [], [], [] + + def layer(index): + a, b = params[index * 2 : index * 2 + 2] + + def compute(value, context): + assert lora_module._CURRENT_LORA_SLOT.get() is owned + calls.append(torch.is_grad_enabled()) + result = _backward_layer(value, a, b) + if torch.is_grad_enabled(): + recomputed.append(weakref.ref(value)) + result.register_hook( + lambda grad: inner_tasks.append( + torch._C._current_graph_task_id() + ) + ) + return result, context + + return compute + + with lora_module.use_lora_slot(owned): + hidden = x + for i in range(2): + hidden, _ = mcore_random.checkpoint(layer(i), False, hidden, None) + head = torch_checkpoint.checkpoint( + lambda z: z.sin().square(), + hidden, + use_reentrant=False, + ) + output, loss = _backward_outputs(hidden, head) + expected = reference_x + for i in range(2): + expected = _backward_layer(expected, *reference_params[i * 2 : i * 2 + 2]) + _, reference_loss = _backward_outputs(expected, expected.sin().square()) + work.attach([output, output]) + assert len(work.outputs) == (4 if enabled else 0) + outer_tasks = [] + for attempt in range(2): + before, event_count = work.work_ns, len(cpu_backward_events) + with lora_module.use_lora_slot(ambient): + loss.backward(retain_graph=attempt == 0) + assert lora_module._CURRENT_LORA_SLOT.get() is ambient + reference_loss.backward(retain_graph=attempt == 0) + for actual, expected_grad in zip( + (x, *params), (reference_x, *reference_params), strict=True + ): + assert actual.grad is not None and expected_grad.grad is not None + torch.testing.assert_close( + actual.grad, expected_grad.grad, rtol=1e-12, atol=1e-12 + ) + assert work.work_ns == before + if enabled: + assert ( + len(work.rows) == 1 and len(cpu_backward_events) == event_count + 1 + ) + task, row = next(iter(work.rows.items())) + assert row.ended is not None and row.ended > row.started + assert not row.blocked and row.tail.task == task + outer_tasks.append(task) + before += row.ended - row.started + else: + assert not work.rows and len(cpu_backward_events) == event_count + work.harvest() + assert work.work_ns == before and not work.rows + work.harvest() + assert work.work_ns == before + assert len(calls) == 6 and sum(calls) == 4 + assert len(inner_tasks) == 4 and not set(inner_tasks).intersection(outer_tasks) + assert all(type(task) is int and task >= 0 for task in inner_tasks) + assert len(set(outer_tasks)) == (2 if enabled else 0) + assert not mcore_random.is_checkpointing() and not work.invalid + assert work.cost_ns > 0 + return recomputed + [ + weakref.ref(t) + for t in ( + x, + *params, + hidden, + head, + output.target_logprobs, + output.logits, + output.top_k.logprobs, + ) + ] + + try: + for _ in range(2): + refs = iteration() + gc.collect() + assert all(ref() is None for ref in refs) + work.attach([]) + assert not work.outputs and not work.rows + finally: + work.close() + owner = weakref.ref(work) + del work + gc.collect() + assert owner() is None + + +def test_backward_work_checkpoint_failure_lifetime(cpu_backward_events): + work = BackwardWork(threading.RLock(), torch.device("cpu")) + work.disabled = False # Test-only CPU readiness facade. + + def failed_iteration(): + original = RuntimeError("checkpoint recomputation failed") + cause = ValueError("original cause") + x, params = _backward_inputs() + owned = lora_module.LoRASlotRef("lora", "owned") + ambient = lora_module.LoRASlotRef("lora", "ambient") + + def compute(value, context): + assert lora_module._CURRENT_LORA_SLOT.get() is owned + if torch.is_grad_enabled(): + raise original from cause + return _backward_layer(value, *params[:2]), context + + with lora_module.use_lora_slot(owned): + hidden, _ = mcore_random.checkpoint(compute, False, x, None) + head = torch_checkpoint.checkpoint( + lambda z: z.sin().square(), + hidden, + use_reentrant=False, + ) + output, loss = _backward_outputs(hidden, head) + work.attach([output]) + with lora_module.use_lora_slot(ambient): + try: + loss.backward() + except RuntimeError as error: + assert error is original and error.__cause__ is cause + else: + pytest.fail("Expected the original recomputation error") + assert lora_module._CURRENT_LORA_SLOT.get() is ambient + assert len(work.rows) == 1 + assert all(row.ended is None and row.tail is None for row in work.rows.values()) + work.harvest() + assert work.work_ns == 0 and not cpu_backward_events + work.close() + assert work.closed and not work.rows and not work.outputs + # This fixture owns the pre-created exception captured by compute. + # Its traceback can retain this frame/graph. Dispose of it only after + # error/cause assertions; production must preserve the original error. + original.__traceback__ = None + assert original.__cause__ is cause + return [ + weakref.ref(t) + for t in ( + x, + *params, + hidden, + head, + output.target_logprobs, + output.logits, + output.top_k.logprobs, + ) + ] + + try: + refs = failed_iteration() + gc.collect() + assert all(ref() is None for ref in refs) + finally: + work.close() + owner = weakref.ref(work) + del work + gc.collect() + assert owner() is None + + @pytest.mark.parametrize("compiled", [False, True]) @pytest.mark.parametrize("pending_graph", [False, True]) def test_dispatcher_cache_releases_checkpoint_inputs( From 59c81a58a6ee584881b73556429cba9bba5f64b0 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 22 Sep 2026 07:46:21 +0000 Subject: [PATCH 6/9] Integrate qualified observer accounting test fixtures --- .../unit/test_trainer_rank_cache_recovery.py | 18 +++++++- .../unit/test_trainer_rank_handoff_budget.py | 42 ++++++++++++++++++- 2 files changed, 58 insertions(+), 2 deletions(-) diff --git a/tests/unit/test_trainer_rank_cache_recovery.py b/tests/unit/test_trainer_rank_cache_recovery.py index 3b6b08f8c..158945a8f 100644 --- a/tests/unit/test_trainer_rank_cache_recovery.py +++ b/tests/unit/test_trainer_rank_cache_recovery.py @@ -8,7 +8,7 @@ import unittest from unittest.mock import patch -from art.trainer_rank import _impl +from art.trainer_rank import _backward_work, _impl Refusal = _impl.TrainerRankMemoryError Partial = _impl.TrainerRankPartialExecutionError @@ -87,6 +87,14 @@ class Clock: def __init__(self): self.value = 0.0 self.next = None + self.nanoseconds = 0 + self.observer_step_ns = 0 + + def perf_counter_ns(self): + # Existing recovery cases isolate O=0; imported-module cases opt in to + # positive observer cost without changing the original episode clock. + self.nanoseconds += self.observer_step_ns + return self.nanoseconds def perf_counter(self): if self.next is not None: @@ -153,6 +161,9 @@ def make(self): patcher = patch.object(_impl, name, value) patcher.start() self.addCleanup(patcher.stop) + observer_clock = patch.object(_backward_work, "time", clock) + observer_clock.start() + self.addCleanup(observer_clock.stop) q = object.__new__(_impl.TrainerRank) q.device = types.SimpleNamespace(type="cuda") q._update_peak_memory_profile = lambda *a: None @@ -574,6 +585,9 @@ def __exit__(self, *args): return False state.lock = Lock() + # This case injects faults by original recovery-lock ordinal; + # observer cancellation/depth has separate actual helper controls. + q._backward_work = lambda: None _, error, _ = run(q, [fail(n), success(n)]) self.assertIs(error, original) self.assertEqual(state.lock.calls, 4) @@ -599,6 +613,7 @@ def __exit__(self, *args): return False state.lock = Lock() + q._backward_work = lambda: None # Isolate original recovery-lock ordinals. _, error, _ = run(q, [fail(n), success(n)]) self.assertIs(error, original) self.assertEqual(state.lock.calls, 4) @@ -620,6 +635,7 @@ def __exit__(self, *args): return False state.lock = Lock() + q._backward_work = lambda: None # Isolate original recovery-lock ordinals. value, error, _ = run(q, [fail(n), success(n)]) self.assertIsNone(error) self.assertTrue(value[1].fits) diff --git a/tests/unit/test_trainer_rank_handoff_budget.py b/tests/unit/test_trainer_rank_handoff_budget.py index 50527bbb7..44476d222 100644 --- a/tests/unit/test_trainer_rank_handoff_budget.py +++ b/tests/unit/test_trainer_rank_handoff_budget.py @@ -7,7 +7,7 @@ import pytest import torch -from art.trainer_rank import ForwardOutput, TrainerRank, _impl +from art.trainer_rank import ForwardOutput, TrainerRank, _backward_work, _impl from tests.unit import test_trainer_rank_cache_recovery as recovery from tests.unit.test_trainer_rank_physical_reserve import allocator from tests.unit.test_trainer_rank_validation import ( @@ -71,6 +71,46 @@ def test_repeat_budget_has_the_same_five_percent_boundary(rig, cost, releases): assert state.cost == cost + 1.0 and state.work == 1000.0 and state.owner is None +@pytest.mark.parametrize( + "observer_step_ns,cost,releases", + [(0, 40.0, 1), (1_000_000, 40.0, 0), (1_000_000, 39.0, 1)], +) +def test_imported_observer_cost_moves_repeat_boundary( + rig, observer_step_ns, cost, releases +): + rank, cuda, clock, _ = rig + assert _impl.BackwardWork is _backward_work.BackwardWork + clock.observer_step_ns = observer_step_ns + state = rank._recovery_state() + state.first_consumed, state.work, state.cost, state.high = True, 1000.0, cost, 10.0 + ticks = iter((1.0, 1.0, 2.0)) + rank._recovery_clock = lambda: next(ticks) + cuda.free = 1 + original_reduce = rank._recovery_reduce + observations = [] + + def reduce(values, *, op, sync_across_dp): + if op == "SUM": + observer = state.backward + assert isinstance(observer, _backward_work.BackwardWork) + observations.append((list(values), observer.cost_ns / 1e9)) + elif op == "MAX" and len(values) == 4: + assert values[1:] == [1000.0, 1.0, 0.0] + return original_reduce(values, op=op, sync_across_dp=sync_across_dp) + + rank._recovery_reduce = reduce + rank._release_cached_memory_for_backward(plan(True)) + assert len(observations) == 1 + operands, observer_seconds = observations[0] + assert operands == [cost + observer_seconds, 10.0] + assert (observer_seconds > 0) == (observer_step_ns > 0) + assert (sum(operands) <= 0.05 * 1000.0) == bool(releases) + assert cuda.events.count("release") == releases + assert state.cost == cost + 1.0 and state.work == 1000.0 + assert state.owner is None and not state.invalid + assert state.backward.work_ns == 0 + + @pytest.mark.parametrize("modes", [(), (False,)]) def test_empty_and_local_no_grad_peer_participate_without_cuda_queries(rig, modes): rank, cuda, _, _ = rig From 5be9bc4c950013e3bab96441555b4d8470c11b85 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 22 Sep 2026 08:01:14 +0000 Subject: [PATCH 7/9] Apply formatter to backward accounting implementation and controls --- src/art/trainer_rank/_backward_work.py | 8 +- src/art/trainer_rank/_impl.py | 3 +- tests/unit/test_trainer_rank_backward_work.py | 220 +++++++++++++----- 3 files changed, 170 insertions(+), 61 deletions(-) diff --git a/src/art/trainer_rank/_backward_work.py b/src/art/trainer_rank/_backward_work.py index 9715b6982..b8d655323 100644 --- a/src/art/trainer_rank/_backward_work.py +++ b/src/art/trainer_rank/_backward_work.py @@ -5,8 +5,8 @@ spans are charged conservatively, including spans also timed by recovery. """ -from dataclasses import dataclass from contextlib import contextmanager, nullcontext +from dataclasses import dataclass from functools import wraps import sys import time @@ -40,6 +40,7 @@ def measured(self, *args, **kwargs): def region(function): """Exclude original forward/recovery work; meter only our transitions.""" + @wraps(function) def wrapped(rank, *args, **kwargs): work = rank._backward_work() @@ -146,7 +147,10 @@ def hook(grad, ref=ref): work._start() # Returning None preserves the original gradient object. - self.outputs[key] = (weakref.ref(tensor), tensor.register_hook(hook)) + self.outputs[key] = ( + weakref.ref(tensor), + tensor.register_hook(hook), + ) @_measured def _start(self): diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 7ff117806..5b6807a9c 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -54,7 +54,8 @@ _local_position_pairs, estimate_prefix_tree_packed_tokens, ) -from art.trainer_rank._backward_work import BackwardWork, region as _backward_region +from art.trainer_rank._backward_work import BackwardWork +from art.trainer_rank._backward_work import region as _backward_region from art.trainer_rank._planner_cost import ( COEFFICIENT_VERSION_FALLBACK, ModelGeometry, diff --git a/tests/unit/test_trainer_rank_backward_work.py b/tests/unit/test_trainer_rank_backward_work.py index 75f7646b5..99d58e6a8 100644 --- a/tests/unit/test_trainer_rank_backward_work.py +++ b/tests/unit/test_trainer_rank_backward_work.py @@ -19,7 +19,6 @@ from unittest.mock import patch import weakref - ROOT = Path(__file__).resolve().parents[2] @@ -102,9 +101,13 @@ def setUp(self): cuda=self.cuda, compiler=NS(is_compiling=lambda: False), _C=NS(_current_graph_task_id=lambda: self.task), - autograd=NS(Variable=NS(_execution_engine=NS( - queue_callback=self.callbacks.append, - ))), + autograd=NS( + Variable=NS( + _execution_engine=NS( + queue_callback=self.callbacks.append, + ) + ) + ), ) name = "_art_backward_work_control" spec = importlib.util.spec_from_file_location( @@ -118,8 +121,12 @@ def setUp(self): self.addCleanup(self.work.close) def output(self, tensor): - return NS(target_logprobs=tensor, logits=tensor, hidden_states=tensor, - top_k=NS(logprobs=tensor)) + return NS( + target_logprobs=tensor, + logits=tensor, + hidden_states=tensor, + top_k=NS(logprobs=tensor), + ) def start(self, task): self.task = task @@ -141,25 +148,60 @@ def completed(self, task=1, *, ready=True): def actual_rank(self): source = ast.parse((ROOT / "src/art/trainer_rank/_impl.py").read_text()) names = {"_CacheRecoveryState", "_MemoryCheck"} - methods = {"_recovery_state", "_backward_work", "_try_cache_recovery", - "_execute_split_plan_with_memory_tracking"} - selected = [node for node in source.body if isinstance(node, ast.ClassDef) - and node.name in names] - trainer = next(node for node in source.body if isinstance(node, ast.ClassDef) - and node.name == "TrainerRank") - trainer.body = [node for node in trainer.body - if isinstance(node, ast.FunctionDef) and node.name in methods] - ns = dict(__name__=__name__, dataclass=dataclass, dataclass_field=field, - threading=threading, BackwardWork=self.module.BackwardWork, - _backward_region=self.module.region, torch=self.torch, - math=math, os=os, time=self.clock, cast=lambda typ, value: value, - TrainerRankMemoryError=type("MemoryRefusal", (RuntimeError,), {}), - _TEST_HOOKS_ENV="ART_BACKWARD_CONTROL_ONLY", - _TEST_MEMORY_LIMIT_ENV="ART_BACKWARD_CONTROL_LIMIT", - _MEMORY_RESERVE_FRACTION=0.05) - tree = ast.Module(body=[ast.ImportFrom(module="__future__", names=[ - ast.alias(name="annotations")], level=0), *selected, trainer], type_ignores=[]) - exec(compile(ast.fix_missing_locations(tree), "actual-accounting-methods", "exec"), ns) + methods = { + "_recovery_state", + "_backward_work", + "_try_cache_recovery", + "_execute_split_plan_with_memory_tracking", + } + selected = [ + node + for node in source.body + if isinstance(node, ast.ClassDef) and node.name in names + ] + trainer = next( + node + for node in source.body + if isinstance(node, ast.ClassDef) and node.name == "TrainerRank" + ) + trainer.body = [ + node + for node in trainer.body + if isinstance(node, ast.FunctionDef) and node.name in methods + ] + ns = dict( + __name__=__name__, + dataclass=dataclass, + dataclass_field=field, + threading=threading, + BackwardWork=self.module.BackwardWork, + _backward_region=self.module.region, + torch=self.torch, + math=math, + os=os, + time=self.clock, + cast=lambda typ, value: value, + TrainerRankMemoryError=type("MemoryRefusal", (RuntimeError,), {}), + _TEST_HOOKS_ENV="ART_BACKWARD_CONTROL_ONLY", + _TEST_MEMORY_LIMIT_ENV="ART_BACKWARD_CONTROL_LIMIT", + _MEMORY_RESERVE_FRACTION=0.05, + ) + tree = ast.Module( + body=[ + ast.ImportFrom( + module="__future__", names=[ast.alias(name="annotations")], level=0 + ), + *selected, + trainer, + ], + type_ignores=[], + ) + exec( + compile( + ast.fix_missing_locations(tree), "actual-accounting-methods", "exec" + ), + ns, + ) rank = object.__new__(ns["TrainerRank"]) rank.device = self.device state = rank._recovery_state() @@ -292,14 +334,19 @@ def fail(rank): self.assertGreater(self.work.cost_ns, 0) self.assertEqual(self.work.depth, 0) - def test_ordinary_queue_fault_does_not_replace_gradient_or_erase_previous_credit(self): + def test_ordinary_queue_fault_does_not_replace_gradient_or_erase_previous_credit( + self, + ): expected = self.completed(1) self.work.harvest() tensor = Tensor(self.device) self.work.attach([self.output(tensor)]) self.task = 2 - with patch.object(self.torch.autograd.Variable._execution_engine, - "queue_callback", side_effect=RuntimeError("queue")): + with patch.object( + self.torch.autograd.Variable._execution_engine, + "queue_callback", + side_effect=RuntimeError("queue"), + ): self.assertIsNone(next(iter(tensor.hooks.values()))(object())) self.assertTrue(self.work.disabled) self.work.harvest() @@ -307,7 +354,12 @@ def test_ordinary_queue_fault_does_not_replace_gradient_or_erase_previous_credit def test_new_cancellation_at_each_observer_boundary_propagates_exact_object(self): original_work = self.work - for error_type in (KeyboardInterrupt, SystemExit, GeneratorExit, CancelledError): + for error_type in ( + KeyboardInterrupt, + SystemExit, + GeneratorExit, + CancelledError, + ): for stage in ("attach", "hook", "callback", "harvest", "leave", "close"): with self.subTest(error=error_type, stage=stage): self.work = self.module.BackwardWork(threading.RLock(), self.device) @@ -339,7 +391,10 @@ def test_new_cancellation_at_each_observer_boundary_propagates_exact_object(self action = lambda: scope.__exit__(None, None, None) else: self.work.attach([self.output(tensor)]) - target, name = next(iter(self.work.outputs.values()))[1], "remove" + target, name = ( + next(iter(self.work.outputs.values()))[1], + "remove", + ) action = self.work.close with patch.object(target, name, side_effect=primary): with self.assertRaises(BaseException) as caught: @@ -357,9 +412,14 @@ def test_meter_propagates_new_cancellation_and_preserves_inflight_primary(self): cause = ValueError("cause") primary.__cause__ = cause self.task = 1 - with patch.object(self.clock, "perf_counter_ns", side_effect=[2000, 2001, secondary]): - with patch.object(self.torch.autograd.Variable._execution_engine, - "queue_callback", side_effect=primary): + with patch.object( + self.clock, "perf_counter_ns", side_effect=[2000, 2001, secondary] + ): + with patch.object( + self.torch.autograd.Variable._execution_engine, + "queue_callback", + side_effect=primary, + ): with self.assertRaises(KeyboardInterrupt) as caught: self.work._start() self.assertIs(caught.exception, primary) @@ -380,7 +440,9 @@ def test_actual_lookup_and_constructor_cancellation_identity(self): primary, secondary = CancelledError("lookup"), GeneratorExit("meter") cause = ValueError("lookup cause") primary.__cause__ = cause - with patch.object(self.clock, "perf_counter_ns", side_effect=[primary, secondary]): + with patch.object( + self.clock, "perf_counter_ns", side_effect=[primary, secondary] + ): with self.assertRaises(CancelledError) as caught: rank._backward_work() self.assertIs(caught.exception, primary) @@ -412,11 +474,16 @@ def body(rank): raise primary from cause return "original result" - with patch.object(self.clock, "perf_counter_ns", - side_effect=[10000, 10001, secondary, 10003]): + with patch.object( + self.clock, + "perf_counter_ns", + side_effect=[10000, 10001, secondary, 10003], + ): with self.assertRaises(BaseException) as caught: body(rank) - self.assertIs(caught.exception, secondary if primary is None else primary) + self.assertIs( + caught.exception, secondary if primary is None else primary + ) self.assertEqual(self.work.depth, 0) if primary is not None: self.assertIs(primary.__cause__, cause) @@ -430,8 +497,11 @@ def untouched(rank): with self.work.region(): self.assertEqual(self.work.depth, 1) - with patch.object(self.clock, "perf_counter_ns", - side_effect=[20000, secondary, 20002, 20003]): + with patch.object( + self.clock, + "perf_counter_ns", + side_effect=[20000, secondary, 20002, 20003], + ): with self.assertRaises(KeyboardInterrupt) as caught: untouched(rank) self.assertIs(caught.exception, secondary) @@ -441,7 +511,9 @@ def untouched(rank): scope.__enter__() self.assertEqual(self.work.depth, 2) another = SystemExit("secondary exit meter") - with patch.object(self.clock, "perf_counter_ns", side_effect=[secondary, another]): + with patch.object( + self.clock, "perf_counter_ns", side_effect=[secondary, another] + ): with self.assertRaises(KeyboardInterrupt) as caught: scope.__exit__(None, None, None) self.assertIs(caught.exception, secondary) @@ -450,7 +522,9 @@ def untouched(rank): def test_clock_fault_cost_overflow_and_unknown_readiness_fail_closed(self): prior = self.work.cost_ns - with patch.object(self.clock, "perf_counter_ns", side_effect=RuntimeError("clock")): + with patch.object( + self.clock, "perf_counter_ns", side_effect=RuntimeError("clock") + ): self.work.harvest() self.assertTrue(self.work.invalid) self.assertEqual(self.work.cost_ns, prior) @@ -472,8 +546,10 @@ def test_wrong_device_mode_id_and_nonbool_readiness_withhold(self): self.assertTrue(self.work.disabled) self.assertFalse(tensor.hooks) self.work.disabled = False - with patch.dict(sys.modules, {"torch._dynamo.compiled_autograd": NS( - compiled_autograd_enabled=True)}): + with patch.dict( + sys.modules, + {"torch._dynamo.compiled_autograd": NS(compiled_autograd_enabled=True)}, + ): self.start(1) self.assertTrue(self.work.disabled) self.work.disabled = False @@ -523,8 +599,12 @@ def reduce(values, *, op, sync_across_dp): return values rank._recovery_reduce = reduce - result = rank._try_cache_recovery(ns["_MemoryCheck"](80, 0, False), - sync_across_dp=True, owner=state.owner, started=0.0) + result = rank._try_cache_recovery( + ns["_MemoryCheck"](80, 0, False), + sync_across_dp=True, + owner=state.owner, + started=0.0, + ) self.assertFalse(result) self.assertEqual([op for op, _ in calls], ["SUM", "MAX", "MIN"]) self.assertEqual(calls[1][1][1], 100.0) @@ -536,21 +616,33 @@ def test_actual_first_repeat_rule_and_invalid_meter(self): rank._recovery_reduce = lambda values, **kw: values state.owner, state.high = object(), 0.1 check = ns["_MemoryCheck"](80, 0, False) - self.assertTrue(rank._try_cache_recovery(check, sync_across_dp=False, - owner=state.owner, started=0.0)) # Original first free release. + self.assertTrue( + rank._try_cache_recovery( + check, sync_across_dp=False, owner=state.owner, started=0.0 + ) + ) # Original first free release. self.cuda.free = 0 - self.assertFalse(rank._try_cache_recovery(check, sync_across_dp=False, - owner=state.owner, started=0.0)) + self.assertFalse( + rank._try_cache_recovery( + check, sync_across_dp=False, owner=state.owner, started=0.0 + ) + ) # Use a genuinely completed host interval in this scalar fixture. self.start(1) self.clock.value += 30_000_000_000 self.finish(1) - self.assertTrue(rank._try_cache_recovery(check, sync_across_dp=False, - owner=state.owner, started=0.0)) + self.assertTrue( + rank._try_cache_recovery( + check, sync_across_dp=False, owner=state.owner, started=0.0 + ) + ) self.cuda.free = 0 self.work.invalid = True - self.assertFalse(rank._try_cache_recovery(check, sync_across_dp=False, - owner=state.owner, started=0.0)) + self.assertFalse( + rank._try_cache_recovery( + check, sync_across_dp=False, owner=state.owner, started=0.0 + ) + ) def test_actual_split_rollback_keeps_b_and_o_and_primary_error(self): rank, state, ns = self.actual_rank() @@ -570,10 +662,16 @@ def execute(*args, **kw): return [object()], None rank._run_flat_plan_with_memory_tracking = execute - plan = NS(request_count=2, subforwards=[object(), object()], - request_indices=[[0], [1]], subforward_count=2) + plan = NS( + request_count=2, + subforwards=[object(), object()], + request_indices=[[0], [1]], + subforward_count=2, + ) with self.assertRaises(ValueError) as caught: - rank._execute_split_plan_with_memory_tracking(plan, check=None, context="test") + rank._execute_split_plan_with_memory_tracking( + plan, check=None, context="test" + ) self.assertIs(caught.exception, primary) self.assertEqual(state.work, 5.0) self.assertEqual(self.work.work_ns, credited) @@ -586,8 +684,14 @@ def test_original_invalid_forward_guard_cannot_be_masked_by_positive_b(self): self.start(1) self.clock.value += 30_000_000_000 self.finish(1) - self.assertFalse(rank._try_cache_recovery(ns["_MemoryCheck"](80, 0, False), - sync_across_dp=False, owner=state.owner, started=0.0)) + self.assertFalse( + rank._try_cache_recovery( + ns["_MemoryCheck"](80, 0, False), + sync_across_dp=False, + owner=state.owner, + started=0.0, + ) + ) self.assertTrue(state.invalid) self.assertEqual(self.cuda.releases, 0) From e8f5a357064ecacd5f076d3c29e3520809349097 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 22 Sep 2026 08:11:50 +0000 Subject: [PATCH 8/9] Clarify completed-row typing and module-loader test setup --- src/art/trainer_rank/_backward_work.py | 3 ++- tests/unit/test_trainer_rank_backward_work.py | 3 ++- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/src/art/trainer_rank/_backward_work.py b/src/art/trainer_rank/_backward_work.py index b8d655323..7f6203709 100644 --- a/src/art/trainer_rank/_backward_work.py +++ b/src/art/trainer_rank/_backward_work.py @@ -284,7 +284,8 @@ def harvest(self): self.disabled = True return if ready and not row.blocked: - addition += row.ended - row.started + # The quiescence check above excludes None under this lock. + addition += row.ended - row.started # ty: ignore[unsupported-operator] if ready or row.blocked: retired.append(task) if self.work_ns + addition > _MAX_NS: diff --git a/tests/unit/test_trainer_rank_backward_work.py b/tests/unit/test_trainer_rank_backward_work.py index 99d58e6a8..e6279eb9f 100644 --- a/tests/unit/test_trainer_rank_backward_work.py +++ b/tests/unit/test_trainer_rank_backward_work.py @@ -113,10 +113,11 @@ def setUp(self): spec = importlib.util.spec_from_file_location( name, ROOT / "src/art/trainer_rank/_backward_work.py" ) + assert spec is not None and spec.loader is not None self.module = importlib.util.module_from_spec(spec) with patch.dict(sys.modules, {"torch": self.torch, name: self.module}): spec.loader.exec_module(self.module) - self.module.time = self.clock + setattr(self.module, "time", self.clock) self.work = self.module.BackwardWork(threading.RLock(), self.device) self.addCleanup(self.work.close) From 5a4ea977552fa9b337ca65d8e3fabfa22fd1e887 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 22 Sep 2026 09:09:51 +0000 Subject: [PATCH 9/9] Preserve forward failures through recovery observer entry --- src/art/trainer_rank/_backward_work.py | 19 ++- src/art/trainer_rank/_impl.py | 112 +++++++----- tests/unit/test_trainer_rank_backward_work.py | 161 ++++++++++++++++++ 3 files changed, 242 insertions(+), 50 deletions(-) diff --git a/src/art/trainer_rank/_backward_work.py b/src/art/trainer_rank/_backward_work.py index 7f6203709..b69b6e9dc 100644 --- a/src/art/trainer_rank/_backward_work.py +++ b/src/art/trainer_rank/_backward_work.py @@ -43,8 +43,14 @@ def region(function): @wraps(function) def wrapped(rank, *args, **kwargs): - work = rank._backward_work() - with work.region() if work is not None else nullcontext(): + error = kwargs.get("error") + try: + work = rank._backward_work() + except BaseException: + if error is None: + raise + work = None + with work.region(error=error) if work is not None else nullcontext(): return function(rank, *args, **kwargs) return wrapped @@ -219,9 +225,9 @@ def _finish(self, task): row.ended, row.tail = ended, tail @contextmanager - def region(self): + def region(self, *, error: BaseException | None = None): """Own one exclusion count, even when a transition clock cancels.""" - started, entered, primary = None, False, False + started, entered, primary = None, False, error is not None try: try: started = time.perf_counter_ns() @@ -238,8 +244,9 @@ def region(self): self.disabled = True except BaseException: self.disabled = True - primary = True - raise + if not primary: + primary = True + raise finally: self._charge(started, primary=primary) yield diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 5b6807a9c..5c0a49c33 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -2463,7 +2463,7 @@ def _release_cached_memory_for_backward( ) -> None: # Every WORLD wave reaches this before the public iterator skips empty # outputs. Forward has already executed: never replan or retry here. - with self._cache_recovery_episode() as (owner, started): + with self._cache_recovery_episode(error=error) as (owner, started): exchange_error: BaseException | None = None try: failed, gradients = self._recovery_reduce( @@ -5281,55 +5281,74 @@ def finish(value: Any) -> Any: raise latest from original @contextmanager - def _cache_recovery_episode(self) -> Iterator[tuple[object, float | None]]: - state = self._recovery_state() - started = self._recovery_clock() - owner = object() - with state.lock: - if state.owner is None: - state.owner = owner - primary: BaseException | None = None + def _cache_recovery_episode( + self, *, error: BaseException | None = None + ) -> Iterator[tuple[object, float | None]]: + state = None + started = None + owner = None + primary = error try: + try: + state = self._recovery_state() + started = self._recovery_clock() + owner = object() + with state.lock: + if state.owner is None: + state.owner = owner + except BaseException: + # A saved forward failure must still reach the WORLD status + # vote. Unknown accounting cost cannot enable later recovery. + if state is None: + state = getattr(self, "_cache_recovery_state", None) + if state is None: + state = self._cache_recovery_state = _CacheRecoveryState() + state.invalid = True + if error is None: + raise yield owner, started except BaseException as exc: primary = exc raise finally: - # Includes control, release, resample and repeated search; no idle. - try: - finished = self._recovery_clock() - elapsed = ( - None if started is None or finished is None else finished - started - ) - with state.lock: - if ( - elapsed is None - or not math.isfinite(elapsed) - or elapsed <= 0 - or not math.isfinite(state.cost + elapsed) - ): - state.invalid = True - else: - state.cost += elapsed - state.high = max(state.high, elapsed) - except Exception: - state.invalid = True - except BaseException as exc: - state.invalid = True - if primary is None: - primary = exc - raise - finally: + if state is not None: + # Includes control, release, resample and repeated search; no idle. try: + finished = self._recovery_clock() + elapsed = ( + None + if started is None or finished is None + else finished - started + ) with state.lock: - if state.owner is owner: - state.owner = None + if ( + elapsed is None + or not math.isfinite(elapsed) + or elapsed <= 0 + or not math.isfinite(state.cost + elapsed) + ): + state.invalid = True + else: + state.cost += elapsed + state.high = max(state.high, elapsed) except Exception: state.invalid = True - except BaseException: + except BaseException as exc: state.invalid = True if primary is None: + primary = exc raise + finally: + try: + with state.lock: + if state.owner is owner: + state.owner = None + except Exception: + state.invalid = True + except BaseException: + state.invalid = True + if primary is None: + raise def _recovery_clock(self) -> float | None: try: @@ -5365,25 +5384,30 @@ def _recovery_state(self) -> _CacheRecoveryState: return state def _backward_work(self) -> BackwardWork | None: - state = self._recovery_state() + state = None started = None primary = False try: + state = self._recovery_state() started = time.perf_counter_ns() with state.lock: if state.backward is None and not state.invalid: state.backward = BackwardWork(state.lock, self.device) return state.backward - except Exception: - # Unknown accounting cost must never become free recovery budget. - state.invalid = True - return None - except BaseException: + except BaseException as exc: + # Unknown accounting cost must never become free recovery budget, + # including a failure before state lookup returned. + if state is None: + state = getattr(self, "_cache_recovery_state", None) + if state is None: + state = self._cache_recovery_state = _CacheRecoveryState() state.invalid = True + if isinstance(exc, Exception): + return None primary = True raise finally: - if state.backward is not None: + if state is not None and state.backward is not None: # Includes lazy setup and lock wait; constructor overlap is an # intentional conservative charge, not an exact subtraction. state.backward._charge(started, primary=primary) diff --git a/tests/unit/test_trainer_rank_backward_work.py b/tests/unit/test_trainer_rank_backward_work.py index e6279eb9f..40ca52ed4 100644 --- a/tests/unit/test_trainer_rank_backward_work.py +++ b/tests/unit/test_trainer_rank_backward_work.py @@ -6,6 +6,7 @@ import ast from asyncio import CancelledError +from contextlib import ExitStack, contextmanager from dataclasses import dataclass, field import gc import importlib.util @@ -14,6 +15,7 @@ from pathlib import Path import sys import threading +import traceback from types import SimpleNamespace as NS import unittest from unittest.mock import patch @@ -153,6 +155,9 @@ def actual_rank(self): "_recovery_state", "_backward_work", "_try_cache_recovery", + "_cache_recovery_episode", + "_release_cached_memory_for_backward", + "_memory_error_with_reduction_note", "_execute_split_plan_with_memory_tracking", } selected = [ @@ -175,6 +180,8 @@ def actual_rank(self): dataclass=dataclass, dataclass_field=field, threading=threading, + contextmanager=contextmanager, + traceback=traceback, BackwardWork=self.module.BackwardWork, _backward_region=self.module.region, torch=self.torch, @@ -521,6 +528,160 @@ def untouched(rank): self.assertEqual(self.work.depth, 1) self.assertEqual(self.work.depth, 0) + def handoff_entry_failure( + self, stage, primary, *, foreign_owner=False, exchange_error=None + ): + self.work = self.module.BackwardWork(threading.RLock(), self.device) + self.addCleanup(self.work.close) + rank, state, _ = self.actual_rank() + state.owner = original_owner = object() if foreign_owner else None + state.work, state.cost = 17.0, 3.0 + self.work.work_ns = 19 + secondary = KeyboardInterrupt(stage) + votes = [] + + def reduce(values, **kwargs): + votes.append((values, kwargs)) + if exchange_error is not None: + raise exchange_error + return values + + rank._recovery_reduce = reduce + clock = self.clock.perf_counter_ns + count = 0 + + def observer_clock(): + nonlocal count + count += 1 + if ( + count + == {"lookup_clock": 1, "region_clock": 3, "region_charge": 4}[stage] + ): + raise secondary + return clock() + + with ExitStack() as patches: + if stage in ("lookup_clock", "region_clock", "region_charge"): + patches.enter_context( + patch.object(self.clock, "perf_counter_ns", observer_clock) + ) + elif stage in ("lookup_state", "episode_state"): + values = ( + [secondary, state] + if stage == "lookup_state" + else [state, secondary] + ) + patches.enter_context( + patch.object(rank, "_recovery_state", side_effect=values) + ) + elif stage == "episode_clock": + patches.enter_context( + patch.object(rank, "_recovery_clock", side_effect=[secondary, 2.0]) + ) + else: + # Cancel after owner assignment, including while a foreign + # episode owns the state: cleanup must never clear its token. + lock = state.lock + cancel_exit = False + + def episode_clock(): + nonlocal cancel_exit + if count == 0: + cancel_exit = True + return 1.0 + + class Lock: + def __enter__(self): + return lock.__enter__() + + def __exit__(self, *args): + nonlocal cancel_exit, count + result = lock.__exit__(*args) + if cancel_exit: + cancel_exit = False + count += 1 + raise secondary + return result + + state.lock = self.work.lock = Lock() + rank._recovery_clock = episode_clock + with self.assertRaises(BaseException) as caught: + rank._release_cached_memory_for_backward( + NS(groups=[NS(grad_enabled=True)]), error=primary + ) + self.assertIs(caught.exception, secondary if primary is None else primary) + self.assertEqual( + votes, + [] + if primary is None + else [([1.0, 1.0], {"op": "MAX", "sync_across_dp": True})], + ) + self.assertIs(state.owner, original_owner) + self.assertEqual(self.work.depth, 0) + self.assertEqual((state.work, self.work.work_ns), (17.0, 19)) + self.assertGreaterEqual(state.cost, 3.0) + self.assertEqual(self.cuda.releases, 0) + self.assertTrue(state.invalid or self.work.invalid or self.work.disabled) + + def test_saved_forward_error_reaches_failure_vote_despite_entry_cancellation(self): + for stage in ( + "lookup_state", + "lookup_clock", + "region_clock", + "region_charge", + "episode_state", + "episode_clock", + "episode_owner", + ): + for error_type in (RuntimeError, CancelledError): + with self.subTest(stage=stage, error_type=error_type): + primary = error_type("forward") + primary.__cause__ = cause = ValueError("original cause") + primary.__context__ = context = LookupError("original context") + self.handoff_entry_failure(stage, primary) + self.assertIs(primary.__cause__, cause) + self.assertIs(primary.__context__, context) + self.assertTrue(primary.__suppress_context__) + self.handoff_entry_failure( + "episode_owner", RuntimeError("forward"), foreign_owner=True + ) + primary = RuntimeError("forward before poisoned exchange") + self.handoff_entry_failure( + "episode_clock", primary, exchange_error=SystemExit("exchange cancellation") + ) + self.assertIn("exchange cancellation", "\n".join(primary.__notes__)) + + def test_new_entry_cancellation_propagates_and_cleans_owned_state(self): + for stage in ( + "lookup_state", + "lookup_clock", + "region_clock", + "region_charge", + "episode_state", + "episode_clock", + "episode_owner", + ): + with self.subTest(stage=stage): + self.handoff_entry_failure(stage, None) + self.handoff_entry_failure("episode_owner", None, foreign_owner=True) + + def test_recovery_episode_still_records_cost_and_preserves_foreign_owner(self): + rank, state, _ = self.actual_rank() + state.cost, state.high = 3.0, 0.25 + for foreign_owner in (None, object()): + with self.subTest(foreign_owner=foreign_owner): + state.owner = foreign_owner + with patch.object(rank, "_recovery_clock", side_effect=[10.0, 10.5]): + with rank._cache_recovery_episode() as (owner, started): + self.assertEqual(started, 10.0) + self.assertIs( + state.owner, + owner if foreign_owner is None else foreign_owner, + ) + self.assertIs(state.owner, foreign_owner) + self.assertFalse(state.invalid) + self.assertEqual((state.cost, state.high), (4.0, 0.5)) + def test_clock_fault_cost_overflow_and_unknown_readiness_fail_closed(self): prior = self.work.cost_ns with patch.object(