You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
fix: skip the CPU memory probe that doubles MultiStartProdigy compile #1646
Nightly release Stage 3 (release-integrate run 36226772178, 2026-09-26) timed out on autolens_workspace/scripts/imaging/start_here.py at 3605 s; the three prior nights took 704-770 s. The cause is MultiStartGradient._warn_if_unbatched_exceeds_memory: whenever batch_size=None it probes Analysis.batched_memory_bytes at batch 1 and 2, and each probe is a full throwaway XLA compile of jit(vmap(value_and_grad(fitness.call))). Its docstring assumes memory_analysis returns 0 on CPU so the probe would bail, but on jax 0.10.2 CPU it is non-zero, so every fresh CPU MultiStart fit pays two extra full-model compiles. Last night one of them drew the slow mode of the bimodal CPU compile (~54 min gap in the log between vis warm-up and "jit compiling the single-point objective").
autofit/non_linear/search/mle/multi_start_gradient/search.py_warn_if_unbatched_exceeds_memory: inside the existing try, import jax and return if jax.default_backend() == "cpu"; then fetch self._memory_budget_bytes() and return if falsy; only then call the batch-1 / batch-2 probes. Keep the swallow-all-exceptions contract.
Rewrite the "Known limitation" docstring paragraph: CPU is skipped because the projection costs two full-model XLA compiles which, on a release-sized lens model, can exceed the fit itself (run 36226772178, imaging/start_here.py 3605 s TIMEOUT, ~54 min probe compile).
test_autofit/non_linear/search/mle/test_multi_start_gradient.py near the batch_memory_model tests: (a) CPU backend -> stub analysis batched_memory_bytes never called (red on unfixed main); (b) jax.default_backend monkeypatched to "gpu" and _memory_budget_bytes to a value -> probe called and the warning path works.
Run the full test_autofit/non_linear/search/mle/ directory.
Key Files
autofit/non_linear/search/mle/multi_start_gradient/search.py — the guard (~line 815, called ~1080)
Nightly release Stage 3 (PyAutoHeart release-integrate run 36226772178) TIMEOUT on autolens_workspace scripts/imaging/start_here.py at 3605s (prior nights 704-770s). Cause: MultiStartGradient._warn_if_unbatched_exceeds_memory (PyAutoFit autofit/non_linear/search/mle/multi_start_gradient/search.py ~line 815, called at ~1080 when batch_size is None) calls analysis.batched_memory_bytes (autofit/non_linear/analysis/analysis.py:338) at batch 1 and 2, each a full throwaway XLA compile of jit(vmap(value_and_grad(fitness.call))). Its docstring assumes memory_analysis returns 0 on CPU so it would bail; on jax 0.10.2 CPU it is non-zero, so both compiles run, then budget from psutil. One drew the slow bimodal compile (~54 min).
Overview
Nightly release Stage 3 (release-integrate run 36226772178, 2026-09-26) timed out on
autolens_workspace/scripts/imaging/start_here.pyat 3605 s; the three prior nights took 704-770 s. The cause isMultiStartGradient._warn_if_unbatched_exceeds_memory: wheneverbatch_size=Noneit probesAnalysis.batched_memory_bytesat batch 1 and 2, and each probe is a full throwaway XLA compile ofjit(vmap(value_and_grad(fitness.call))). Its docstring assumesmemory_analysisreturns 0 on CPU so the probe would bail, but on jax 0.10.2 CPU it is non-zero, so every fresh CPU MultiStart fit pays two extra full-model compiles. Last night one of them drew the slow mode of the bimodal CPU compile (~54 min gap in the log between vis warm-up and "jit compiling the single-point objective").Plan
BUILD_SCRIPT_TIMEOUToverride in autolens_workspace stays for now.Detailed implementation plan
Work Classification
Library
Affected Repositories
Branch Survey
Suggested branch:
feature/multistart-cpu-memory-probeWorktree root:
~/Code/PyAutoLabs-wt/multistart-cpu-memory-probe/Implementation Steps
autofit/non_linear/search/mle/multi_start_gradient/search.py_warn_if_unbatched_exceeds_memory: inside the existing try,import jaxand return ifjax.default_backend() == "cpu"; then fetchself._memory_budget_bytes()and return if falsy; only then call the batch-1 / batch-2 probes. Keep the swallow-all-exceptions contract.test_autofit/non_linear/search/mle/test_multi_start_gradient.pynear thebatch_memory_modeltests: (a) CPU backend -> stub analysisbatched_memory_bytesnever called (red on unfixed main); (b)jax.default_backendmonkeypatched to "gpu" and_memory_budget_bytesto a value -> probe called and the warning path works.test_autofit/non_linear/search/mle/directory.Key Files
autofit/non_linear/search/mle/multi_start_gradient/search.py— the guard (~line 815, called ~1080)autofit/non_linear/analysis/analysis.py—batched_memory_bytes(~line 338)test_autofit/non_linear/search/mle/test_multi_start_gradient.py— testsOriginal Prompt
Click to expand starting prompt
Nightly release Stage 3 (PyAutoHeart release-integrate run 36226772178) TIMEOUT on autolens_workspace scripts/imaging/start_here.py at 3605s (prior nights 704-770s). Cause:
MultiStartGradient._warn_if_unbatched_exceeds_memory(PyAutoFitautofit/non_linear/search/mle/multi_start_gradient/search.py~line 815, called at ~1080 when batch_size is None) callsanalysis.batched_memory_bytes(autofit/non_linear/analysis/analysis.py:338) at batch 1 and 2, each a full throwaway XLA compile of jit(vmap(value_and_grad(fitness.call))). Its docstring assumes memory_analysis returns 0 on CPU so it would bail; on jax 0.10.2 CPU it is non-zero, so both compiles run, then budget from psutil. One drew the slow bimodal compile (~54 min).