Skip to content

fix: skip the CPU memory probe that doubles MultiStartProdigy compile #1646

Description

@Jammy2211

Overview

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").

Plan

  • Skip the memory projection entirely on the CPU backend, before any probe compile, restoring the documented intent (the guard exists for GPU OOM, fix: bound MultiStartProdigy vmap memory (interferometer release-leg OOM) #1452).
  • Check the memory budget before probing, so an undeterminable budget never pays for a compile.
  • Rewrite the docstring's "Known limitation" paragraph to say CPU is skipped deliberately and why, citing run 36226772178.
  • Add unit tests: CPU never calls the probe; a (monkeypatched) GPU backend still does.
  • No workspace change; the 3600 s BUILD_SCRIPT_TIMEOUT override in autolens_workspace stays for now.
Detailed implementation plan

Work Classification

Library

Affected Repositories

  • PyAutoFit (primary)

Branch Survey

Repository Current Branch Dirty?
./PyAutoFit main clean

Suggested branch: feature/multistart-cpu-memory-probe
Worktree root: ~/Code/PyAutoLabs-wt/multistart-cpu-memory-probe/

Implementation Steps

  1. 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.
  2. 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).
  3. 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.
  4. 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)
  • autofit/non_linear/analysis/analysis.py — batched_memory_bytes (~line 338)
  • test_autofit/non_linear/search/mle/test_multi_start_gradient.py — tests

Original 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 (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).

No activity

Activity on this issue will appear here.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions