From 3220566e356e5a461f22b0c20db08e03070ef648 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Sat, 26 Sep 2026 11:14:36 +0100 Subject: [PATCH] fix: skip the CPU memory probe that doubles MultiStart compile MultiStartGradient._warn_if_unbatched_exceeds_memory probed Analysis.batched_memory_bytes at batch 1 and 2 whenever batch_size=None, and each probe is a full throwaway XLA compile of the batched value_and_grad. It relied on memory_analysis returning 0 on CPU to bail after the first probe; from jax 0.10.2 CPU reports a non-zero value, so every CPU MultiStart fit paid two extra full-model compiles. One drew the slow compile mode (~54 min) and pushed autolens imaging/start_here.py to a 3605 s TIMEOUT in release-integrate run 36226772178. Return before any probe on the CPU backend (the guard exists for GPU OOM, #1452), and read the memory budget before probing so an unknown budget never pays a compile. Tests pin both gates with stub probes. Closes #1646 Co-Authored-By: Claude Opus 5.5 --- .../search/mle/multi_start_gradient/search.py | 34 ++++++--- .../search/mle/test_multi_start_gradient.py | 70 +++++++++++++++++++ 2 files changed, 93 insertions(+), 11 deletions(-) diff --git a/autofit/non_linear/search/mle/multi_start_gradient/search.py b/autofit/non_linear/search/mle/multi_start_gradient/search.py index c24a20d02..374f269df 100644 --- a/autofit/non_linear/search/mle/multi_start_gradient/search.py +++ b/autofit/non_linear/search/mle/multi_start_gradient/search.py @@ -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. @@ -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 diff --git a/test_autofit/non_linear/search/mle/test_multi_start_gradient.py b/test_autofit/non_linear/search/mle/test_multi_start_gradient.py index 7caaa9b87..06b965b40 100644 --- a/test_autofit/non_linear/search/mle/test_multi_start_gradient.py +++ b/test_autofit/non_linear/search/mle/test_multi_start_gradient.py @@ -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