Skip to content
Open
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
1 change: 1 addition & 0 deletions src/art/trainer_rank/_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
37 changes: 28 additions & 9 deletions src/art/trainer_rank/_micro_batch_planner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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,
Expand All @@ -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()
Expand Down
Loading
Loading