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 4ee194dfb..38aeb12d4 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, @@ -5187,14 +5234,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() @@ -5221,6 +5261,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 @@ -5335,11 +5387,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() @@ -5371,7 +5424,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), @@ -5388,16 +5441,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( @@ -5427,7 +5484,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 @@ -5445,9 +5502,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 @@ -5467,7 +5552,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)