diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index a58663c33..84096a7b1 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -663,6 +663,7 @@ class _CandidateMicroBatch(Generic[ForwardInputsT]): rejected_candidates: int cold_start: bool fallback: _CandidateMicroBatch[ForwardInputsT] | None = None + recovery_probe: _MemoryCheck | None = None class _SlotGraphSentinel(torch.autograd.Function): diff --git a/src/art/trainer_rank/_micro_batch_planner.py b/src/art/trainer_rank/_micro_batch_planner.py index 2a077b57d..0f495a272 100644 --- a/src/art/trainer_rank/_micro_batch_planner.py +++ b/src/art/trainer_rank/_micro_batch_planner.py @@ -929,10 +929,12 @@ def candidate(width: int) -> _CandidateMicroBatch[ForwardInputsT]: try: outcome = self._admission_outcome( 0 - if admission_error is not None + if found is None else 1 if isinstance(found, _impl._ForwardRefusal) else 2 + if found[0].subforward_count > 1 + else 3 ) except BaseException: if admission_error is not None: @@ -961,6 +963,9 @@ def candidate(width: int) -> _CandidateMicroBatch[ForwardInputsT]: stats_global_count=min_width, rejected_candidates=len(rejected_widths), cold_start=True, + # The existing WORLD vote makes this uniform even on peers + # with a flat or empty share. Retain demand, not a second plan. + recovery_probe=first.check if outcome == 2 else None, ) if first.cold_start: return first @@ -1828,7 +1833,7 @@ def discard_planner_observation(self: TrainerRank) -> None: def _admission_outcome(self: TrainerRank, local: int) -> int: - """Existing world fallback MIN: error=0, refusal=1, fit=2.""" + """Existing world fallback MIN: error=0, refusal=1, split=2, flat=3.""" if not (_impl.dist.is_available() and _impl.dist.is_initialized()): return local value = _impl.torch.tensor( @@ -1968,7 +1973,7 @@ def _recover_admission_impl( sync_across_dp: bool, admit_refusal: Callable[[_ForwardRefusal], Any] | None = None, ) -> Any: - """Pure search, at most one smaller-plan refresh, then one recovery.""" + """Pure search, one budgeted recovery, and freshly checked fallback.""" original: TrainerRankMemoryError | None = None refused: _ForwardRefusal | None = None best: _ForwardRefusal | None = None @@ -2067,12 +2072,18 @@ def finish(value: Any) -> Any: value = search() result = finish(value) - if result is not None: + probe = ( + result.recovery_probe + if isinstance(result, _impl._CandidateMicroBatch) + else None + ) + if result is not None and probe is None: return result - assert refused is not None - original = refused.error(context) + incumbent = result + if refused is not None: + original = refused.error(context) with self._cache_recovery_episode() as (owner, started): - if not isinstance(value, _impl._ForwardRefusal): + if incumbent is None and not isinstance(value, _impl._ForwardRefusal): # A formerly fitting width is not proof that the minimum cannot fit. value = search() result = finish(value) @@ -2083,9 +2094,11 @@ def finish(value: Any) -> Any: assert refused is not None self._snapshot_planning_telemetry(refused.plan, refused.check) return reject() - assert refused is not None + if probe is None: + assert refused is not None + probe = refused.check if self._try_cache_recovery( - refused.check, + probe, sync_across_dp=sync_across_dp, owner=owner, started=started, @@ -2094,6 +2107,12 @@ def finish(value: Any) -> Any: result = finish(value) if result is not None: return result + if incumbent is not None: + # Denial or ineffective release is harmless only while the saved + # split still fits fresh WORLD counters. Never attempt release twice. + result = finish(incumbent) + if result is not None: + return result assert refused is not None self._snapshot_planning_telemetry(refused.plan, refused.check) return reject() diff --git a/tests/unit/test_trainer_rank_cache_upgrade.py b/tests/unit/test_trainer_rank_cache_upgrade.py new file mode 100644 index 000000000..45c810abb --- /dev/null +++ b/tests/unit/test_trainer_rank_cache_upgrade.py @@ -0,0 +1,310 @@ +"""One minimum-wave cache upgrade, using actual search/recovery and CPU facades.""" + +from dataclasses import replace +from types import SimpleNamespace +from typing import cast + +import pytest +from test_trainer_rank_handoff_budget import plan as handoff_plan +from test_trainer_rank_handoff_budget import rig +from test_trainer_rank_split import _packed_budget, _rank, _request + +from art.trainer_rank import _impl + + +def candidate(required=5, *, probe=True): + return _impl._CandidateMicroBatch( + inputs=[], + indices=(0,), + plan=cast( + _impl._AnyForwardPlan, + SimpleNamespace(packed_tokens=required, logical_tokens=required), + ), + check=_impl._MemoryCheck(required, 10, required <= 10), + stats_global_count=1, + rejected_candidates=1, + cold_start=True, + recovery_probe=_impl._MemoryCheck(80, 10, False) if probe else None, + ) + + +def admit(rank, values, calls): + def search(): + value = values[min(len(calls), len(values) - 1)] + calls.append(value) + if isinstance(value, BaseException): + raise value + return value + + return rank._recover_admission( + search, + lambda value: (value.plan, value.check), + lambda value, check: replace(value, check=check), + context="forward_micro_batches", + sync_across_dp=True, + ) + + +def test_upgrade_first_attempt_rebuilds_search_and_shares_handoff_budget(rig): + rank, cuda, _, _ = rig + calls = [] + flat = candidate(80, probe=False) + result = admit(rank, [candidate(), flat], calls) + assert result.plan is flat.plan and result.check.available_bytes == 170 + assert len(calls) == 2 and cuda.events.count("release") == 1 + state = rank._recovery_state() + assert state.first_consumed and state.cost > 0 and state.owner is None + cuda.free = 1 + rank._release_cached_memory_for_backward(handoff_plan(True)) + assert cuda.events.count("release") == 1 + rank._record_recovery_work("forward_micro_batches", 100.0) + rank._release_cached_memory_for_backward(handoff_plan(True)) + assert cuda.events.count("release") == 2 + + +@pytest.mark.parametrize("cost,releases", [(40.0, 1), (40.5, 0)]) +def test_upgrade_uses_existing_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) + calls = [] + result = admit(rank, [candidate(), candidate(80, probe=False)], calls) + assert result.check.estimated_required_bytes == (80 if releases else 5) + assert cuda.events.count("release") == releases and len(calls) == 1 + releases + assert state.cost == cost + 1.0 and state.owner is None + + +@pytest.mark.parametrize("searched", ["split", "refused"]) +def test_ineffective_upgrade_returns_fresh_split_without_second_attempt(rig, searched): + rank, cuda, _, ns = rig + cuda.empty_cache = lambda: cuda.events.append("release") + incumbent = candidate() + retry = ( + incumbent + if searched == "split" + else ns["_ForwardRefusal"]( + incumbent.plan, _impl._MemoryCheck(80, 10, False), "still refused" + ) + ) + calls = [] + result = admit(rank, [incumbent, retry], calls) + assert result.plan is incumbent.plan and result.check.available_bytes == 10 + assert len(calls) == 2 and cuda.events.count("release") == 1 + assert rank._recovery_state().owner is None + + +@pytest.mark.parametrize("release", [False, True]) +def test_stale_incumbent_refuses_without_second_release(rig, release): + rank, cuda, _, _ = rig + if release: + + def empty_cache(): + cuda.events.append("release") + cuda.free = 31 + + cuda.empty_cache = empty_cache + else: + + def denied(*args, **kwargs): + cuda.free = 31 + return False + + rank._try_cache_recovery = denied + calls = [] + with pytest.raises(_impl.TrainerRankMemoryError): + admit(rank, [candidate()], calls) + assert cuda.events.count("release") == int(release) + assert len(calls) == 1 + int(release) and rank._recovery_state().owner is None + + +@pytest.mark.parametrize("stage", ["release", "search"]) +def test_upgrade_preserves_original_failure(rig, stage): + rank, cuda, _, _ = rig + error = RuntimeError(stage) + if stage == "release": + cuda.failure = error + with pytest.raises(RuntimeError) as caught: + admit(rank, [candidate(), error], []) + assert caught.value is error + assert cuda.events.count("release") == 1 and rank._recovery_state().owner is None + + +def test_normal_fit_has_no_upgrade_episode(rig): + rank, _, _, _ = rig + rank._cache_recovery_episode = lambda: pytest.fail("no recovery metadata") + assert admit(rank, [candidate(probe=False)], []).check.fits + + +def test_actual_minimum_wave_search_upgrades_only_after_physical_release(monkeypatch): + rank = _rank(monkeypatch) + requests = [_request(i) for i in range(4)] + monkeypatch.setattr(rank, "_retained_memory_bytes", lambda *a, **k: 0) + free, total = [50], 1000 + _packed_budget(monkeypatch, rank, lambda: free[0] - 30) + monkeypatch.setattr(rank, "device", _impl.torch.device("cuda")) + monkeypatch.setattr(_impl.torch.cuda, "is_available", lambda: True) + monkeypatch.setattr(_impl.torch.cuda, "get_allocator_backend", lambda: "native") + monkeypatch.setattr( + _impl.torch.cuda, "mem_get_info", lambda device: (free[0], total) + ) + monkeypatch.setattr(_impl.torch.cuda, "memory_allocated", lambda device: 0) + monkeypatch.delenv(_impl._TEST_HOOKS_ENV, raising=False) + releases, searches = [], [] + search = rank._search_next_micro_batch + + def observed(*args, **kwargs): + result = search(*args, **kwargs) + searches.append(result) + return result + + def release(): + releases.append(free[0]) + free[0] = 90 + + monkeypatch.setattr(rank, "_search_next_micro_batch", observed) + monkeypatch.setattr(_impl.torch.cuda, "empty_cache", release) + selected = rank._select_next_micro_batch([requests], 0) + assert len(searches) == 2 and searches[0].plan.subforward_count == 2 + assert searches[0].recovery_probe.estimated_required_bytes == 40 + assert selected.plan.subforward_count == 1 and selected.plan.packed_tokens == 40 + assert selected.check.available_bytes == 60 and releases == [50] + assert not _impl.torch.cuda.is_initialized() + + +@pytest.mark.parametrize( + "local,outcome", + [("split", 2), ("flat", 2), ("empty", 2), ("flat", 3), ("empty", 3)], +) +def test_fallback_world_status_sets_uniform_probe_before_fit_return( + monkeypatch, local, outcome +): + rank = _rank(monkeypatch) + requests = [_request(i) for i in range(4)] + monkeypatch.setattr(rank, "_retained_memory_bytes", lambda *a, **k: 0) + _packed_budget(monkeypatch, rank, 20) + if local == "split": + found = rank._find_admissible_forward( + requests, checkpoint=_impl.Unset, refusal_prefix="test" + ) + else: + plan = rank._plan_flat_forward([] if local == "empty" else requests[:1]) + found = plan, rank._memory_check(plan) + monkeypatch.setattr(rank, "_dp_rank_and_size", lambda: (int(local == "empty"), 2)) + monkeypatch.setattr(rank, "_estimate_flat_forward", lambda *a, **k: None) + monkeypatch.setattr( + rank, "_memory_check", lambda *a, **k: _impl._MemoryCheck(40, 20, False) + ) + monkeypatch.setattr(rank, "_find_admissible_forward", lambda *a, **k: found) + events = [] + + def vote(status): + events.append(("status", status)) + return outcome + + def refresh(check, **kw): + events.append(("refresh",)) + return replace( + check, available_bytes=20, fits=check.estimated_required_bytes <= 20 + ) + + def reduce(values, *, op, sync_across_dp): + assert sync_across_dp + events.append((op,)) + return values + + monkeypatch.setattr(rank, "_admission_outcome", vote) + monkeypatch.setattr(rank, "_refresh_memory_check", refresh) + monkeypatch.setattr(rank, "_recovery_reduce", reduce) + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 20) + selected = rank._select_next_micro_batch([requests], 0) + assert (selected.recovery_probe is not None) == (outcome == 2) + assert events == [("status", 2 if local == "split" else 3), ("refresh",)] + ( + [("SUM",), ("MAX",), ("MIN",), ("refresh",)] if outcome == 2 else [] + ) + assert selected.inputs == ([] if local == "empty" else [requests]) + + +def test_handoff_first_attempt_is_not_renewed_for_upgrade(rig): + rank, cuda, _, _ = rig + cuda.free = 1 + rank._release_cached_memory_for_backward(handoff_plan(True)) + cuda.free = 40 + calls = [] + result = admit(rank, [candidate()], calls) + assert result.check.fits and len(calls) == 1 + assert cuda.events.count("release") == 1 + + +@pytest.mark.parametrize("reason", ["backend", "cap", "foreign_owner"]) +def test_upgrade_preserves_native_cap_and_owner_stops(rig, monkeypatch, reason): + rank, cuda, _, _ = rig + if reason == "backend": + cuda.backend = "cudaMallocAsync" + elif reason == "cap": + monkeypatch.setenv("CONTROL_ART_HOOK", "1") + monkeypatch.setenv("CONTROL_ART_LIMIT", "40") + else: + rank._recovery_state().owner = object() + calls = [] + assert admit(rank, [candidate()], calls).check.fits + # Existing non-native availability includes cached bytes, so it can cause + # a fresh search without any native release. This policy is unchanged. + assert len(calls) == (2 if reason == "backend" else 1) + assert "release" not in cuda.events + + +@pytest.mark.parametrize("outcome", [0, 1]) +def test_failed_fallback_vote_never_creates_upgrade_probe(monkeypatch, outcome): + rank = _rank(monkeypatch) + monkeypatch.setattr(rank, "_retained_memory_bytes", lambda *a, **k: 0) + _packed_budget(monkeypatch, rank, 20) + statuses = [] + monkeypatch.setattr( + rank, "_admission_outcome", lambda status: statuses.append(status) or outcome + ) + if outcome == 0: + with pytest.raises(RuntimeError, match="another DP rank"): + rank._search_next_micro_batch([[_request(i) for i in range(4)]], 0) + else: + result = rank._search_next_micro_batch([[_request(i) for i in range(4)]], 0) + assert isinstance(result, _impl._ForwardRefusal) + assert statuses == [2] + + +def test_local_fallback_error_precedes_vote_error(monkeypatch): + rank = _rank(monkeypatch) + _packed_budget(monkeypatch, rank, 20) + primary = RuntimeError("local planner") + + def failed(*args, **kwargs): + raise primary + + def vote(status): + assert status == 0 + raise ValueError("secondary exchange") + + monkeypatch.setattr(rank, "_find_admissible_forward", failed) + monkeypatch.setattr(rank, "_admission_outcome", vote) + with pytest.raises(RuntimeError) as caught: + rank._search_next_micro_batch([[_request(i) for i in range(4)]], 0) + assert caught.value is primary + + +def test_trust_only_cold_plan_does_not_offer_cache_upgrade(monkeypatch): + rank = _rank(monkeypatch) + _packed_budget(monkeypatch, rank, 100) + monkeypatch.setattr(rank, "_all_ranks_have_memory_profile", lambda **kw: False) + monkeypatch.setattr( + rank, + "_find_admissible_forward", + lambda *a, **k: pytest.fail("no memory rejection"), + ) + monkeypatch.setattr( + rank, + "_cache_recovery_episode", + lambda: pytest.fail("trust does not enable release"), + ) + selected = rank._select_next_micro_batch([[_request(i) for i in range(4)]], 0) + assert selected.cold_start and selected.recovery_probe is None diff --git a/tests/unit/test_trainer_rank_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py index 21b7ad071..de9777b95 100644 --- a/tests/unit/test_trainer_rank_head_memory.py +++ b/tests/unit/test_trainer_rank_head_memory.py @@ -93,7 +93,7 @@ def test_recovery_keeps_profile_demand_above_cold_head_floor(monkeypatch): ) -def test_real_head_split_fits_before_cache_recovery(monkeypatch): +def test_real_head_split_survives_denied_cache_upgrade(monkeypatch): from test_trainer_rank_split import _recording_executor from art.trainer_rank import TrainerRank, _impl @@ -124,8 +124,9 @@ def test_real_head_split_fits_before_cache_recovery(monkeypatch): monkeypatch.setattr( r, "_available_memory_bytes", lambda: TrainerRank._available_memory_bytes(probe) ) + probes = [] monkeypatch.setattr( - r, "_try_cache_recovery", lambda *a, **kw: pytest.fail("Split already fits") + r, "_try_cache_recovery", lambda check, **kw: probes.append(check) or False ) executed = _recording_executor(monkeypatch, r) batches = list(r.forward_micro_batches([requests])) @@ -136,6 +137,7 @@ def test_real_head_split_fits_before_cache_recovery(monkeypatch): for group in batches[0].outputs ] == [[7, 107, 207, 307]] assert r.last_forward_telemetry()["subforward_request_indices"] == ((0, 1), (2, 3)) + assert len(probes) == 1 and probes[0].estimated_required_bytes > budget assert not torch.cuda.is_initialized() diff --git a/tests/unit/test_trainer_rank_split_peak.py b/tests/unit/test_trainer_rank_split_peak.py index 4ed9aff00..598071080 100644 --- a/tests/unit/test_trainer_rank_split_peak.py +++ b/tests/unit/test_trainer_rank_split_peak.py @@ -272,10 +272,14 @@ def test_partial_forward_does_not_learn_split_peak(monkeypatch): assert rank._memory_profiles == counters["profiles_before_failure"] -def test_empty_dp_rank_retains_global_selection_collective_sequence(monkeypatch): - # Real selection/find/rung methods with explicit scalar reductions. This is - # not native distributed convergence or a model-execution test. +@pytest.mark.parametrize("request_count", [1, 2]) +def test_empty_dp_rank_retains_global_selection_collective_sequence( + monkeypatch, request_count +): + # Real selection/find/rung/recovery methods with simulated peer operands. + # This is not native distributed convergence or a model-execution test. traces = [] + split = request_count == 2 for dp_rank in (0, 1): with monkeypatch.context() as patch: rank = _rank() @@ -316,29 +320,91 @@ def searched(*args, **kwargs): patch.setattr(rank, "_search_next_micro_batch", searched) def reduce(value, op, group=None): - trace.append(("global" if group is None else "local", str(op))) - if group is None: - value.fill_( - max(value.item(), 100 if search_finished else 200) - if op == tr.dist.ReduceOp.MAX - else min(value.item(), 100) + trace.append( + ( + "global" if group is None else "local", + str(op), + tuple(value.shape), + value.dtype, ) + ) + if group is None: + if value.dtype == torch.int32: + # WORLD fallback MIN: split=2, flat/empty=3. The empty + # peer must see the split vote and join recovery too. + assert value.ndim == 0 and op == tr.dist.ReduceOp.MIN + assert value.item() == (2 if split and dp_rank == 0 else 3) + peer = value.new_tensor(2 if split else 3) + elif value.ndim == 0: + # Scalar admission checks exchange byte counters only. + assert value.dtype == torch.float64 + assert op in (tr.dist.ReduceOp.MAX, tr.dist.ReduceOp.MIN) + peer = value.new_tensor( + (100 if search_finished else 200) + if op == tr.dist.ReduceOp.MAX + else 100 + ) + else: + # Recovery vectors mix seconds, bytes and flags. Use + # elementwise peer operands, never a scalar byte clamp. + assert split and search_finished + assert value.dtype == torch.float64 + peer = value.new_tensor( + [0.0, 0.0] + if op == tr.dist.ReduceOp.SUM + else [200.0, 0.0, 0.0, 0.0] + if op == tr.dist.ReduceOp.MAX + else [100.0, 0.0, 1.0, 1.0] + ) + assert value.shape == peer.shape + if op == tr.dist.ReduceOp.SUM: + value.add_(peer) + elif op == tr.dist.ReduceOp.MAX: + value.copy_(torch.maximum(value, peer)) + else: + assert op == tr.dist.ReduceOp.MIN + value.copy_(torch.minimum(value, peer)) patch.setattr(tr.dist, "is_available", lambda: True) patch.setattr(tr.dist, "is_initialized", lambda: True) patch.setattr(tr.dist, "all_reduce", reduce) - candidate = rank._select_next_micro_batch([_requests()], 0) + candidate = rank._select_next_micro_batch([_requests(request_count)], 0) assert candidate.check.fits + assert candidate.check.estimated_required_bytes == 100 + assert candidate.check.available_bytes == 100 assert len(candidate.inputs) == (1 if dp_rank == 0 else 0) + assert candidate.indices == ((0,) if dp_rank == 0 else ()) + assert candidate.stats_global_count == 1 assert isinstance( candidate.plan, - tr._SplitForwardPlan if dp_rank == 0 else tr._FlatForwardPlan, + tr._SplitForwardPlan if split and dp_rank == 0 else tr._FlatForwardPlan, ) + assert candidate.plan.request_count == ( + request_count if dp_rank == 0 else 0 + ) + assert len(search_finished) == 1 + assert (candidate.recovery_probe is not None) == split + if split: + assert candidate.recovery_probe.estimated_required_bytes == 200 + state = rank._recovery_state() + assert state.owner is None and not state.first_consumed traces.append([event for event in trace if event[0] == "global"]) assert traces[-1][-2:] == [ - ("global", str(tr.dist.ReduceOp.MAX)), - ("global", str(tr.dist.ReduceOp.MIN)), + ("global", str(tr.dist.ReduceOp.MAX), (), torch.float64), + ("global", str(tr.dist.ReduceOp.MIN), (), torch.float64), ] + vectors = [event for event in traces[-1] if len(event) == 4 and event[2]] + assert vectors == ( + [ + ("global", str(tr.dist.ReduceOp.SUM), (2,), torch.float64), + ("global", str(tr.dist.ReduceOp.MAX), (4,), torch.float64), + ("global", str(tr.dist.ReduceOp.MIN), (4,), torch.float64), + ] + if split + else [] + ) + if split: + assert traces[-1][-5:-2] == vectors assert traces[0] == traces[1]