Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
71 changes: 68 additions & 3 deletions src/art/trainer_rank/_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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.

Expand All @@ -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)
Expand All @@ -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(
Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand All @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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(
Expand Down
15 changes: 11 additions & 4 deletions tests/unit/test_trainer_rank_head_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 [
Expand Down
59 changes: 57 additions & 2 deletions tests/unit/test_trainer_rank_split_peak.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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.
Expand All @@ -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)
Expand All @@ -328,13 +367,29 @@ 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(
candidate.plan,
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)),
Expand Down
Loading
Loading