Skip to content
Merged
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
34 changes: 23 additions & 11 deletions autofit/non_linear/search/mle/multi_start_gradient/search.py
Original file line number Diff line number Diff line change
Expand Up @@ -828,12 +828,17 @@ def _warn_if_unbatched_exceeds_memory(self, model, analysis):
changing the execution shape of every existing run on the strength of a
projection is not a trade this should make on its own.

**Known limitation.** ``memory_analysis`` reports 0 on a CPU-only JAX
build, which is where the release harness runs. The projection is
skipped entirely in that case rather than guessing, so this guard
currently helps GPU users and leaves the CPU path to the (now
unfiltered) traceback. Making the CPU path measurable is follow-up
work and needs validating against a real run, not a unit test.
**CPU is skipped deliberately.** Each probe is a full, throwaway XLA
compile of the batched ``value_and_grad`` (at batch 1 and batch 2), and
on a release-sized lens model those two compiles can cost more than
the fit itself. This was once skipped by accident, because
``memory_analysis`` reported 0 on CPU; from jax 0.10.2 it reports a
non-zero value, both compiles ran, and one of them drew the slow CPU
compile mode (~54 min), pushing ``imaging/start_here.py`` to a 3605 s
TIMEOUT in release-integrate run 36226772178 (2026-09-26). The guard
exists for device-memory OOM on GPU (PyAutoFit#1452), so the CPU path
is left to the (unfiltered) traceback. The memory budget is also read
before any probe, so an undeterminable budget never pays a compile.

Any failure here is swallowed: a memory projection must never be the
reason a fit does not start.
Expand All @@ -843,18 +848,25 @@ def _warn_if_unbatched_exceeds_memory(self, model, analysis):
if probe is None or self.n_starts is None or self.n_starts <= 1:
return

bytes_at_1 = probe(model=model, batch_size=1, gradient=True)
if not bytes_at_1: # None (not on JAX) or 0 (CPU-only: unmeasurable)
return
import jax

bytes_at_2 = probe(model=model, batch_size=2, gradient=True)
if not bytes_at_2:
# Each probe below is a full model compile; never pay for them on
# CPU (see the docstring).
if jax.default_backend() == "cpu":
return

budget = self._memory_budget_bytes()
if not budget:
return

bytes_at_1 = probe(model=model, batch_size=1, gradient=True)
if not bytes_at_1: # None (not on JAX) or 0 (unmeasurable)
return

bytes_at_2 = probe(model=model, batch_size=2, gradient=True)
if not bytes_at_2:
return

fixed, per_start = batch_memory_model(bytes_at_1, bytes_at_2)
suggested = batch_size_within_budget(
fixed, per_start, self.n_starts, budget
Expand Down
70 changes: 70 additions & 0 deletions test_autofit/non_linear/search/mle/test_multi_start_gradient.py
Original file line number Diff line number Diff line change
Expand Up @@ -1077,6 +1077,76 @@ def test__batch_size_within_budget__degenerate_inputs(
assert batch_size_within_budget(fixed, per_start, n_starts, budget) == expected


# The guard that consumes the projection. Each `batched_memory_bytes` probe is
# a full throwaway XLA compile of the batched value_and_grad, so the guard's
# backend gate is what decides whether a fit pays for two of them. The stubs
# below record calls instead of compiling; `jax.default_backend` is patched so
# the tests pin the gate rather than whichever backend the suite runs on.


class _RecordingProbeAnalysis:
def __init__(self, bytes_at_1=2 * GB + INCIDENT_PER_START):
self.calls = []
self._bytes_at_1 = bytes_at_1

def batched_memory_bytes(self, model, batch_size, gradient):
self.calls.append(batch_size)
return self._bytes_at_1 + (batch_size - 1) * INCIDENT_PER_START


class _RecordingLogger:
def __init__(self):
self.warnings = []

def warning(self, message):
self.warnings.append(message)


def test__unbatched_memory_guard__never_probes_on_cpu(monkeypatch):
# Release-integrate run 36226772178: on CPU the two probe compiles ran in
# full (jax 0.10.2 CPU reports non-zero memory_analysis) and one drew the
# slow compile mode, pushing imaging/start_here.py past 3600 s.
jax = pytest.importorskip("jax")
monkeypatch.setattr(jax, "default_backend", lambda: "cpu")

search = af.MultiStartProdigy(n_starts=48)
monkeypatch.setattr(search, "_memory_budget_bytes", lambda: 16 * GB)
analysis = _RecordingProbeAnalysis()

search._warn_if_unbatched_exceeds_memory(model=None, analysis=analysis)

assert analysis.calls == []


def test__unbatched_memory_guard__probes_and_warns_on_gpu(monkeypatch):
jax = pytest.importorskip("jax")
monkeypatch.setattr(jax, "default_backend", lambda: "gpu")

search = af.MultiStartProdigy(n_starts=48)
monkeypatch.setattr(search, "_memory_budget_bytes", lambda: 16 * GB)
search._logger = _RecordingLogger()
analysis = _RecordingProbeAnalysis()

search._warn_if_unbatched_exceeds_memory(model=None, analysis=analysis)

assert analysis.calls == [1, 2]
assert len(search._logger.warnings) == 1
assert "batch_size=" in search._logger.warnings[0]


def test__unbatched_memory_guard__unknown_budget_skips_the_probe(monkeypatch):
jax = pytest.importorskip("jax")
monkeypatch.setattr(jax, "default_backend", lambda: "gpu")

search = af.MultiStartProdigy(n_starts=48)
monkeypatch.setattr(search, "_memory_budget_bytes", lambda: None)
analysis = _RecordingProbeAnalysis()

search._warn_if_unbatched_exceeds_memory(model=None, analysis=analysis)

assert analysis.calls == []


# --- Per-lane best preservation (PyAutoFit#1514) ------------------------------
#
# Same contract as the _nan_lane_counts tests above: the update rule is pure
Expand Down
Loading