diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index d4718cab0..11c262ad0 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -626,6 +626,7 @@ class _CandidateMicroBatch(Generic[ForwardInputsT]): rejected_candidates: int cold_start: bool fallback: _CandidateMicroBatch[ForwardInputsT] | None = None + recovery_target: tuple["_FlatForwardPlan", _MemoryCheck] | None = None class _SlotGraphSentinel(torch.autograd.Function): @@ -3148,6 +3149,7 @@ def _find_admissible_forward( checkpoint: AdapterSelection, refusal_prefix: str, ensure_slots: bool = True, + unsplit_targets: list[tuple[_FlatForwardPlan, _MemoryCheck]] | None = None, ) -> tuple[_AnyForwardPlan, _MemoryCheck] | _ForwardRefusal: """Find an admissible plan: unsplit first, then the bounded split ladder. @@ -3172,6 +3174,8 @@ def _find_admissible_forward( ) check = self._memory_check(plan) if check.fits: + if unsplit_targets is not None: + unsplit_targets[:] = [(plan, check)] return plan, check best = (plan, check) # Best effort before splitting: the memory-minimal (full sharing) @@ -3181,9 +3185,15 @@ def _find_admissible_forward( ) check = self._memory_check(plan) if check.fits: + if unsplit_targets is not None: + unsplit_targets[:] = [(plan, check)] return plan, check if check.estimated_required_bytes < best[1].estimated_required_bytes: best = (plan, check) + if unsplit_targets is not None: + # Preserve the exact already-priced unsplit operand before the + # ladder can replace best with a smaller internal split. + unsplit_targets[:] = [best] request_count = len(requests) if request_count == 1: return _ForwardRefusal( @@ -5380,11 +5390,13 @@ def candidate(width: int) -> _CandidateMicroBatch[ForwardInputsT]: "smallest DP microbatch is predicted to exceed available memory" ) admission_error: BaseException | None = None + unsplit_targets: list[tuple[_FlatForwardPlan, _MemoryCheck]] = [] try: found = self._find_admissible_forward( list(_flatten(local_inputs)), checkpoint=checkpoint, refusal_prefix=refusal_prefix, + unsplit_targets=unsplit_targets, ) except BaseException as exc: admission_error, found = exc, None @@ -5394,6 +5406,11 @@ def candidate(width: int) -> _CandidateMicroBatch[ForwardInputsT]: if admission_error is not None else 1 if isinstance(found, _ForwardRefusal) + else 3 + if isinstance( + cast("tuple[_AnyForwardPlan, _MemoryCheck]", found)[0], + _FlatForwardPlan, + ) else 2 ) except BaseException: @@ -5415,6 +5432,10 @@ def candidate(width: int) -> _CandidateMicroBatch[ForwardInputsT]: "to find a feasible split for its share", ) split_plan, split_check = found + # The existing MIN outcome carries the split bit too. Flat and + # empty local shares retain their exact target when any peer + # splits; pure search never releases cache itself. + recovery_target = unsplit_targets[0] if outcome == 2 else None return _CandidateMicroBatch( inputs=local_inputs, indices=indices, @@ -5423,6 +5444,7 @@ def candidate(width: int) -> _CandidateMicroBatch[ForwardInputsT]: stats_global_count=min_width, rejected_candidates=len(rejected_widths), cold_start=True, + recovery_target=recovery_target, ) if first.cold_start: return first @@ -6945,7 +6967,7 @@ def _memory_check( return self._memory_check_required(required, sync_across_dp=sync_across_dp) def _admission_outcome(self, local: int) -> int: - """Existing world fallback MIN: error=0, refusal=1, fit=2.""" + """World fallback MIN: error=0, refusal=1, split fit=2, flat fit=3.""" if not (dist.is_available() and dist.is_initialized()): return local value = torch.tensor( @@ -7122,6 +7144,8 @@ def reject() -> Any: decision = _planner_evidence.current(self) if decision is not None: decision.outcome = "admitted_oversized" + if isinstance(selected, _CandidateMicroBatch): + selected = replace(selected, recovery_target=None) return selected decision = _planner_evidence.current(self) if decision is not None: @@ -7154,6 +7178,8 @@ def finish(value: Any) -> Any: if sync_across_dp: check = self._refresh_memory_check(check, sync_across_dp=True) if check.fits: + if isinstance(value, _CandidateMicroBatch): + value = replace(value, recovery_target=None) return update(value, check) refused = _ForwardRefusal( plan, @@ -7184,6 +7210,44 @@ def finish(value: Any) -> Any: value = search() result = finish(value) if result is not None: + target = ( + value.recovery_target + if isinstance(value, _CandidateMicroBatch) + else None + ) + if target is not None: + with self._cache_recovery_episode() as (owner, started): + _planner_evidence.record( + self, + "recovery", + "before_internal_split", + target_required_bytes=target[1].estimated_required_bytes, + target_sample_ordinal=( + None + if target[1].sample is None + else target[1].sample.ordinal + ), + target_packed_tokens=target[0].packed_tokens, + ) + if self._try_cache_recovery( + target[1], + sync_across_dp=sync_across_dp, + owner=owner, + started=started, + require_unused_cache=True, + ): + result = finish(search()) + else: + # Optional recovery may decline; preserve the fitting + # split, refreshing its original operand, not its price. + result = finish(result) + if result is None: + result = finish(search()) + if result is not None: + return result + assert refused is not None + self._snapshot_planning_telemetry(refused.plan, refused.check) + return reject() # Never a second release in this admission. return result assert refused is not None original = refused.error(context) @@ -7403,6 +7467,7 @@ def _try_cache_recovery( owner: object, started: float | None, handoff_grad: bool = False, + require_unused_cache: bool = False, ) -> bool: state = self._recovery_state() backward = self._backward_work() @@ -7492,11 +7557,11 @@ def _try_cache_recovery( ): free, total = torch.cuda.mem_get_info(self.device) needed = int(free) < required + int(total * _MEMORY_RESERVE_FRACTION) - if check is None: + if check is None or require_unused_cache: needed &= int(torch.cuda.memory_reserved(self.device)) > int( torch.cuda.memory_allocated(self.device) ) - elif os.environ.get(_TEST_HOOKS_ENV) == "1": + if check is not None and os.environ.get(_TEST_HOOKS_ENV) == "1": limit = os.environ.get(_TEST_MEMORY_LIMIT_ENV) if limit: cap_blocks = required > max( diff --git a/tests/unit/test_trainer_rank_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py index 21b7ad071..03ecf4a74 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_fitting_split_after_cache_recovery_denial(monkeypatch): from test_trainer_rank_split import _recording_executor from art.trainer_rank import TrainerRank, _impl @@ -124,11 +124,18 @@ def test_real_head_split_fits_before_cache_recovery(monkeypatch): monkeypatch.setattr( r, "_available_memory_bytes", lambda: TrainerRank._available_memory_bytes(probe) ) - monkeypatch.setattr( - r, "_try_cache_recovery", lambda *a, **kw: pytest.fail("Split already fits") - ) + denied_check = r._memory_check(whole) + assert not denied_check.fits + recovery_attempts = [] + + def deny_recovery(check, *, require_unused_cache=False, **kwargs): + recovery_attempts.append((check, require_unused_cache)) + return False + + monkeypatch.setattr(r, "_try_cache_recovery", deny_recovery) executed = _recording_executor(monkeypatch, r) batches = list(r.forward_micro_batches([requests])) + assert recovery_attempts == [(denied_check, True)] assert len(batches) == 1 and batches[0].stats.subforward_count == 2 assert batches[0].stats.global_count == 1 and len(executed) == 2 assert [ diff --git a/tests/unit/test_trainer_rank_split_peak.py b/tests/unit/test_trainer_rank_split_peak.py index 4ed9aff00..4442af17a 100644 --- a/tests/unit/test_trainer_rank_split_peak.py +++ b/tests/unit/test_trainer_rank_split_peak.py @@ -273,8 +273,8 @@ def test_partial_forward_does_not_learn_split_peak(monkeypatch): 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. + # Real selection/find/rung/recovery methods with explicit peer reductions. + # CPU recovery declines without a release; this is not native convergence. traces = [] for dp_rank in (0, 1): with monkeypatch.context() as patch: @@ -291,6 +291,9 @@ def test_empty_dp_rank_retains_global_selection_collective_sequence(monkeypatch) ) patch.setattr(rank, "_retained_memory_bytes", lambda *a, **k: 0) patch.setattr(rank, "_ensure_checkpoint_slots_for", lambda *a, **k: None) + patch.setattr( + torch.cuda, "empty_cache", lambda: pytest.fail("CPU recovery released") + ) # Mode/profile agreements are already collective in production; # retain their sequence while controlling this two-rank witness. @@ -314,9 +317,45 @@ def searched(*args, **kwargs): return result patch.setattr(rank, "_search_next_micro_batch", searched) + recovery_reductions = [] def reduce(value, op, group=None): trace.append(("global" if group is None else "local", str(op))) + if value.ndim: + if group is not None: + assert group == f"tp-cp-{dp_rank}" and not search_finished + assert op == tr.dist.ReduceOp.MIN + assert value.dtype == torch.float64 and value.tolist() == [0] + return # Existing local oversized-admission guard. + assert group is None and search_finished + assert value.dtype == torch.float64 + expected = [ + (tr.dist.ReduceOp.SUM, (2,)), + (tr.dist.ReduceOp.MAX, (4,)), + (tr.dist.ReduceOp.MIN, (4,)), + ] + assert len(recovery_reductions) < len(expected) + assert (op, tuple(value.shape)) == expected[ + len(recovery_reductions) + ] + recovery_reductions.append((op, tuple(value.shape))) + if op == tr.dist.ReduceOp.SUM: + assert torch.isfinite(value).all() and (value >= 0).all() + assert value[1].item() == 0 # No earlier recovery high-water. + value.mul_(2) # Model the peer's equal accounting cost. + elif op == tr.dist.ReduceOp.MAX: + assert value.tolist() == [200 if dp_rank == 0 else 0, 0, 0, 0] + value[0] = 200 # The empty peer retains the global deficit. + else: + assert value.tolist() == [100, 0, 1, 1] + return + if value.dtype == torch.int32: + assert group is None and op == tr.dist.ReduceOp.MIN + assert value.item() == (2 if dp_rank == 0 else 3) + value.fill_(2) # A peer splits even when this local rank is empty. + return + assert value.dtype == torch.float64 + assert op in (tr.dist.ReduceOp.MAX, tr.dist.ReduceOp.MIN) if group is None: value.fill_( max(value.item(), 100 if search_finished else 200) @@ -328,6 +367,10 @@ def reduce(value, op, group=None): patch.setattr(tr.dist, "is_initialized", lambda: True) patch.setattr(tr.dist, "all_reduce", reduce) candidate = rank._select_next_micro_batch([_requests()], 0) + assert len(search_finished) == 1 + assert len(recovery_reductions) == 3 + assert candidate.recovery_target is None + assert not rank._recovery_state().first_consumed assert candidate.check.fits assert len(candidate.inputs) == (1 if dp_rank == 0 else 0) assert isinstance( @@ -335,6 +378,18 @@ def reduce(value, op, group=None): tr._SplitForwardPlan if dp_rank == 0 else tr._FlatForwardPlan, ) traces.append([event for event in trace if event[0] == "global"]) + assert traces[-1][-7:] == [ + ("global", str(op)) + for op in ( + tr.dist.ReduceOp.MAX, + tr.dist.ReduceOp.MIN, + tr.dist.ReduceOp.SUM, + tr.dist.ReduceOp.MAX, + tr.dist.ReduceOp.MIN, + tr.dist.ReduceOp.MAX, + tr.dist.ReduceOp.MIN, + ) + ] assert traces[-1][-2:] == [ ("global", str(tr.dist.ReduceOp.MAX)), ("global", str(tr.dist.ReduceOp.MIN)), diff --git a/tests/unit/test_trainer_rank_split_recovery.py b/tests/unit/test_trainer_rank_split_recovery.py new file mode 100644 index 000000000..987faa420 --- /dev/null +++ b/tests/unit/test_trainer_rank_split_recovery.py @@ -0,0 +1,464 @@ +"""Test-first: one existing-budget recovery before a fitting minimum-wave split. + +No real allocator, process group, model, or GPU is used. Real planner integration +uses tiny CPU inputs; recovery uses the existing scalar clock/allocator facade. +""" + +from dataclasses import replace +import gc +import types +from typing import cast +import unittest +import weakref + +import pytest +import test_trainer_rank_cache_recovery as recovery_fixture +from test_trainer_rank_split import _packed_budget, _rank, _request + +from art.trainer_rank import _impl + + +class TestSplitRecovery(unittest.TestCase): + make = recovery_fixture.TestRecovery.make + + def candidate(self, required=5, *, split=True, target=True): + plan = types.SimpleNamespace( + packed_tokens=40, logical_tokens=40, subforward_count=2 if split else 1 + ) + check = _impl._MemoryCheck(required, 10, required <= 10) + value = _impl._CandidateMicroBatch( + inputs=[["A", "B"]], + indices=(0,), + plan=cast(_impl._AnyForwardPlan, plan), + check=check, + stats_global_count=1, + rejected_candidates=1, + cold_start=True, + ) + original = types.SimpleNamespace(packed_tokens=40, logical_tokens=40) + failed = _impl._MemoryCheck(80, 10, False) + # Inject the private contract on the old tree so red tests demonstrate + # ignored recovery, rather than stopping at a constructor TypeError. + object.__setattr__( + value, "recovery_target", (original, failed) if target else None + ) + return value + + def recover(self, rank, candidates): + calls = [] + + def search(): + value = candidates[min(len(calls), len(candidates) - 1)] + calls.append(value) + if isinstance(value, BaseException): + raise value + return value + + result = 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, + ) + return result, calls + + def setup_recovery(self): + q, c, k, n = self.make() + c.memory_reserved = lambda device: 128 + return q, c, k, n + + def test_exact_target_once_then_fresh_search_and_selected_check(self): + q, c, k, n = self.setup_recovery() + first = self.candidate() + selected = self.candidate(80, split=False, target=False) + seen = [] + original = q._try_cache_recovery + + def observed(check, **kwargs): + seen.append((check, kwargs)) + return original(check, **kwargs) + + q._try_cache_recovery = observed + result, calls = self.recover(q, [first, selected]) + self.assertEqual(c.events.count("release"), 1) + self.assertEqual(len(calls), 2) + self.assertIs(seen[0][0], first.recovery_target[1]) + self.assertTrue(seen[0][1]["require_unused_cache"]) + self.assertIs(result.plan, selected.plan) + self.assertEqual(result.check.estimated_required_bytes, 80) + self.assertEqual(result.check.available_bytes, 170) + self.assertIsNone(result.recovery_target) + + def test_budget_denial_preserves_fitting_split(self): + q, c, k, n = self.setup_recovery() + state = q._recovery_state() + state.first_consumed = True + state.cost, state.high, state.work = 1.0, 1.0, 1.0 + first = self.candidate() + result, calls = self.recover(q, [first]) + self.assertIs(result.plan, first.plan) + self.assertIsNone(result.recovery_target) + self.assertEqual(len(calls), 1) + self.assertNotIn("release", c.events) + self.assertGreater(state.cost, 1.0) + self.assertIsNone(state.owner) + + def test_no_unused_cache_preserves_split_without_release(self): + q, c, k, n = self.setup_recovery() + c.memory_reserved = lambda device: c.memory_allocated(device) + first = self.candidate() + result, calls = self.recover(q, [first]) + self.assertIs(result.plan, first.plan) + self.assertIsNone(result.recovery_target) + self.assertEqual(len(calls), 1) + self.assertNotIn("release", c.events) + self.assertFalse(q._recovery_state().first_consumed) + + def test_existing_budget_allows_later_recovery_without_new_allowance(self): + q, c, k, n = self.setup_recovery() + state = q._recovery_state() + state.first_consumed = True + state.cost, state.high, state.work = 0.01, 0.01, 100.0 + result, calls = self.recover( + q, [self.candidate(), self.candidate(80, target=False)] + ) + self.assertEqual(c.events.count("release"), 1) + self.assertEqual(len(calls), 2) + self.assertGreater(state.cost, 0.01) + self.assertTrue(state.first_consumed) + + def test_unsupported_allocator_or_invalid_budget_keeps_fitting_split(self): + for condition in ("backend", "invalid_budget"): + with self.subTest(condition=condition): + q, c, k, n = self.setup_recovery() + if condition == "backend": + c.backend = "cudaMallocAsync" + c.memory_reserved = lambda device: 16 + else: + q._recovery_state().invalid = True + first = self.candidate() + result, calls = self.recover(q, [first]) + self.assertIs(result.plan, first.plan) + self.assertIsNone(result.recovery_target) + self.assertEqual(len(calls), 1) + self.assertNotIn("release", c.events) + self.assertFalse(q._recovery_state().first_consumed) + + def test_no_hint_means_no_optional_recovery(self): + for split in (False, True): + with self.subTest(split=split): + q, c, k, n = self.setup_recovery() + first = self.candidate(split=split, target=False) + q._try_cache_recovery = lambda *a, **kw: self.fail( + "unrequested recovery" + ) + result, calls = self.recover(q, [first]) + self.assertIs(result.plan, first.plan) + self.assertEqual(len(calls), 1) + + def test_insufficient_release_can_keep_new_fitting_split(self): + q, c, k, n = self.setup_recovery() + first, later = self.candidate(), self.candidate(8) + result, calls = self.recover(q, [first, later]) + self.assertEqual(c.events.count("release"), 1) + self.assertEqual(len(calls), 2) + self.assertIs(result.plan, later.plan) + self.assertEqual(result.check.estimated_required_bytes, 8) + + def test_failed_research_never_gets_second_release(self): + q, c, k, n = self.setup_recovery() + refusal = _impl._ForwardRefusal( + cast( + _impl._AnyForwardPlan, + types.SimpleNamespace(packed_tokens=200, logical_tokens=200), + ), + _impl._MemoryCheck(200, 170, False), + "still cannot fit", + ) + with self.assertRaises(_impl.TrainerRankMemoryError): + self.recover(q, [self.candidate(), refusal]) + self.assertEqual(c.events.count("release"), 1) + self.assertIsNone(q._recovery_state().owner) + + def test_oversized_return_releases_recovery_target_after_world_refresh_fails(self): + q, c, k, n = self.setup_recovery() + q._allow_oversized_batches = True + available = iter((10, 3)) + + def refresh(check, **kwargs): + free = next(available) + return replace( + check, available_bytes=free, fits=check.estimated_required_bytes <= free + ) + + q._refresh_memory_check = refresh + + class Target: + packed_tokens = 40 + + first, later = self.candidate(), self.candidate(6) + later = replace(later, recovery_target=(Target(), later.recovery_target[1])) + target = weakref.ref(later.recovery_target[0]) + plan, inputs, indices = later.plan, later.inputs, later.indices + pending = [first, later] + result = q._recover_admission( + lambda: pending.pop(0), + lambda value: (value.plan, value.check), + lambda value, check: replace(value, check=check), + context="forward_micro_batches", + sync_across_dp=True, + admit_refusal=lambda refusal: self.fail("existing candidate must be used"), + ) + self.assertEqual(pending, []) + self.assertEqual(c.events.count("release"), 1) + self.assertIs(result.plan, plan) + self.assertIs(result.inputs, inputs) + self.assertIs(result.indices, indices) + self.assertEqual(result.check, _impl._MemoryCheck(6, 3, False)) + self.assertIsNone(result.recovery_target) + self.assertIsNotNone(later.recovery_target) + del first, later + gc.collect() + self.assertIsNone(target()) + + def test_release_error_preserves_primary_instead_of_running_split(self): + q, c, k, n = self.setup_recovery() + primary = RuntimeError("release failed") + c.failure = primary + with self.assertRaises(RuntimeError) as caught: + self.recover(q, [self.candidate()]) + self.assertIs(caught.exception, primary) + self.assertEqual(c.events.count("release"), 1) + self.assertTrue(q._recovery_state().first_consumed) + self.assertIsNone(q._recovery_state().owner) + + def test_research_error_keeps_primary_and_owner_closes(self): + q, c, k, n = self.setup_recovery() + primary = ValueError("planner failure") + with self.assertRaises(ValueError) as caught: + self.recover(q, [self.candidate(), primary]) + self.assertIs(caught.exception, primary) + self.assertEqual(c.events.count("release"), 1) + self.assertIsNone(q._recovery_state().owner) + + def test_stale_fallback_after_budget_denial_searches_without_release(self): + q, c, k, n = self.setup_recovery() + state = q._recovery_state() + state.first_consumed = True + state.cost, state.high, state.work = 1.0, 1.0, 1.0 + original = q._try_cache_recovery + + def deny(check, **kwargs): + result = original(check, **kwargs) + self.assertFalse(result) + c.free = 33 # Available 3, original split requires 5. + return result + + q._try_cache_recovery = deny + later = self.candidate(2, target=False) + result, calls = self.recover(q, [self.candidate(), later]) + self.assertIs(result.plan, later.plan) + self.assertEqual(len(calls), 2) + self.assertNotIn("release", c.events) + + def test_asymmetric_peers_follow_same_recovery_collective_order(self): + transcripts = [] + for deficient in (False, True): + with self.subTest(deficient=deficient): + q, c, k, n = self.setup_recovery() + c.free = 40 if deficient else 200 + q._refresh_memory_check = lambda check, **kw: replace( + check, + available_bytes=170, + fits=check.estimated_required_bytes <= 170, + ) + expected = [ + ("SUM", 2, [0.02, 0.0]), + ("MAX", 4, [80, 100, 0, 0]), + ("MIN", 4, [10, -1, 1, 1]), + ("MIN", 2, [170, -1]), + ] + transcript = [] + + def reduce(values, *, op, sync_across_dp): + self.assertTrue(sync_across_dp) + want_op, length, result = expected[len(transcript)] + self.assertEqual((op, len(values)), (want_op, length)) + transcript.append((op, len(values))) + return result + + q._recovery_reduce = reduce + result, calls = self.recover( + q, [self.candidate(), self.candidate(80, split=False, target=False)] + ) + self.assertEqual(len(calls), 2) + self.assertEqual(c.events.count("release"), int(deficient)) + self.assertEqual(len(transcript), 4) + transcripts.append(transcript) + self.assertEqual(*transcripts) + + +def test_pure_finder_retains_exact_unsplit_plan_and_check( + monkeypatch: pytest.MonkeyPatch, +): + rank = _rank(monkeypatch) + monkeypatch.setattr(rank, "_retained_memory_bytes", lambda *_a, **_kw: 0) + _packed_budget(monkeypatch, rank, 20) + observed = [] + original = rank._memory_check + + def check(plan, **kwargs): + value = original(plan, **kwargs) + observed.append((plan, value)) + return value + + monkeypatch.setattr(rank, "_memory_check", check) + monkeypatch.setattr( + rank, "_try_cache_recovery", lambda *a, **kw: pytest.fail("impure search") + ) + targets = [] + found = rank._find_admissible_forward( + [_request(i) for i in range(4)], + checkpoint=_impl.Unset, + refusal_prefix="test", + unsplit_targets=targets, + ) + assert not isinstance(found, _impl._ForwardRefusal) + plan, selected = found + assert isinstance(plan, _impl._SplitForwardPlan) + assert len(targets) == 1 + target, denied = targets[0] + assert isinstance(target, _impl._FlatForwardPlan) + assert not denied.fits and selected.fits + assert any(p is target and c is denied for p, c in observed) + assert denied is not selected + + +def test_minimum_wave_hint_does_not_broaden_dp_rank_forward( + monkeypatch: pytest.MonkeyPatch, +): + rank = _rank(monkeypatch) + monkeypatch.setattr(rank, "_retained_memory_bytes", lambda *_a, **_kw: 0) + _packed_budget(monkeypatch, rank, 20) + inputs = [_request(i) for i in range(4)] + candidate = rank._search_next_micro_batch([inputs], 0) + assert isinstance(candidate, _impl._CandidateMicroBatch) + assert isinstance(candidate.plan, _impl._SplitForwardPlan) + assert candidate.stats_global_count == 1 + assert candidate.recovery_target is not None + assert not candidate.recovery_target[1].fits + # The ordinary DP-local return remains its existing (plan, check) pair. + found = rank._find_admissible_forward( + inputs, checkpoint=_impl.Unset, refusal_prefix="test" + ) + assert isinstance(found, tuple) and len(found) == 2 + + +def test_profile_only_minimum_wave_has_no_recovery_hint( + monkeypatch: pytest.MonkeyPatch, +): + rank = _rank(monkeypatch) + _packed_budget(monkeypatch, rank, 1000) + monkeypatch.setattr(rank, "_all_ranks_have_memory_profile", lambda **kw: False) + candidate = rank._search_next_micro_batch([[_request(0), _request(1)]], 0) + assert isinstance(candidate, _impl._CandidateMicroBatch) + assert isinstance(candidate.plan, _impl._FlatForwardPlan) + assert candidate.recovery_target is None + + +@pytest.mark.parametrize("peer", (0, 1, 2)) +def test_search_healthy_peer_carries_target_when_other_peer_splits(monkeypatch, peer): + rank = _rank(monkeypatch) + monkeypatch.setattr(rank, "_dp_rank_and_size", lambda: (peer, 3)) + monkeypatch.setattr(rank, "_retained_memory_bytes", lambda *a, **kw: 0) + _packed_budget(monkeypatch, rank, 20 if peer == 1 else 1000) + local_check = rank._memory_check_required + + def check(required, *, sync_across_dp=False): + return ( + _impl._MemoryCheck(40, 20, False) + if sync_across_dp + else local_check(required) + ) + + monkeypatch.setattr(rank, "_memory_check_required", check) + events = [] + + def outcome(local): + events.append(("outcome", local)) + return 2 + + def all_true(local): + assert not events, "extra vote after the existing outcome collective" + return local + + monkeypatch.setattr(rank, "_admission_outcome", outcome) + monkeypatch.setattr(rank, "_all_ranks_true", all_true) + monkeypatch.setattr( + rank, "_try_cache_recovery", lambda *a, **kw: pytest.fail("impure search") + ) + inputs = [[_request(i) for i in range(2)], [_request(i + 2) for i in range(4)]] + candidate = rank._search_next_micro_batch(inputs, 0) + assert isinstance(candidate, _impl._CandidateMicroBatch) + assert candidate.indices == ((peer,) if peer < 2 else ()) + assert candidate.stats_global_count == 2 + assert isinstance(candidate.plan, _impl._FlatForwardPlan) == (peer != 1) + assert candidate.recovery_target is not None + assert candidate.recovery_target[1].fits == (peer != 1) + assert events == [("outcome", 2 if peer == 1 else 3)] + + +@pytest.mark.parametrize("outcome", (0, 1)) +def test_peer_error_or_refusal_precedes_new_split_vote(monkeypatch, outcome): + rank = _rank(monkeypatch) + monkeypatch.setattr(rank, "_retained_memory_bytes", lambda *a, **kw: 0) + _packed_budget(monkeypatch, rank, 20) + voted = [] + + def exchange(local): + assert local == 2 + voted.append(local) + return outcome + + def all_true(local): + assert not voted, "new split vote after collective failure/refusal" + return local + + monkeypatch.setattr(rank, "_admission_outcome", exchange) + monkeypatch.setattr(rank, "_all_ranks_true", all_true) + inputs = [[_request(i) for i in range(4)]] + if outcome == 0: + with pytest.raises(RuntimeError, match="another DP rank"): + rank._search_next_micro_batch(inputs, 0) + else: + result = rank._search_next_micro_batch(inputs, 0) + assert isinstance(result, _impl._ForwardRefusal) + assert voted == [2] + + +def test_existing_outcome_all_flat_does_not_request_recovery(monkeypatch): + rank = _rank(monkeypatch) + _packed_budget(monkeypatch, rank, 1000) + local_check = rank._memory_check_required + + def check(required, *, sync_across_dp=False): + return ( + _impl._MemoryCheck(40, 20, False) + if sync_across_dp + else local_check(required) + ) + + monkeypatch.setattr(rank, "_memory_check_required", check) + outcomes = [] + + def outcome(local): + outcomes.append(local) + return local + + monkeypatch.setattr(rank, "_admission_outcome", outcome) + candidate = rank._search_next_micro_batch([[_request(i) for i in range(4)]], 0) + assert isinstance(candidate, _impl._CandidateMicroBatch) + assert outcomes == [3] + assert candidate.recovery_target is None + assert isinstance(candidate.plan, _impl._FlatForwardPlan)