Skip to content

refactor: Fitness.objective factory, PoolFactory + JAX fork rule, start_points and the run(ctx) bridge (search-extensibility A2) - #1679

Merged
Jammy2211 merged 6 commits into
mainfrom
feature/search-ext-a2-objective-bridge
Oct 8, 2026
Merged

Jammy2211 merged 6 commits into
mainfrom
feature/search-ext-a2-objective-bridge

Conversation

@Jammy2211

Copy link
Copy Markdown
Collaborator

Summary

Phase A2 of the search-extensibility epic (PyAutoFit#1676). Fitness.objective(kind) is the one lazily-jitted objective factory (scalar, batched, value_and_grad, batched_value_and_grad; per-kind cache; compile=False escape hatch), ported to every call site except NSS (A3b): Emcee, Zeus, Drawer and the initializer now run jitted under JAX instead of eagerly; BFGS/LBFGS on JAX pass jac= and stop finite-differencing; Nautilus and SMC take the batched objective; NUTS the scalar log-density. make_fitness(analysis, model, **overrides) builds from the declared objective_target/invalid_value; use_jax_jit/use_jax_vmap become deprecated aliases (FutureWarning when passed; defaults now None). parallel.PoolFactory owns the EP core guard and the one JAX fork rule across search, grid/sensitivity and analysis pools; start_points(model, fitness, n) plots the start once; the minimal FitContext/run(ctx) bridge ships with Drawer and Nautilus migrated as proofs (every other _fit keeps working). Six commits in the plan's order. --auto run at effective level safe; tier judge → human /prm. A3 (#1677) stacks on this branch; A3b on both.

Flagged items for the reviewer

  1. Design note revision. docs/design/run_ctx.md was frozen at A1; the bridge needed six member-level clarifications, recorded as "Revision 1" (signature and member list unchanged): ctx.schedule.next_budget(done, total)/chunks(total) take the total budget; ctx.start_points(n) returns arrays (Drawer opts out of the plot via _plots_start_point = False); ctx.update hands the live object, synchronously; ctx.resume/ctx.checkpointer wired by A3; ctx.rng is a SeedSequence of the search seed; ctx.close(failed) added. Plus search-side hooks fitness_overrides(analysis) (Nautilus's batched/batch_size) and samples_cls. Rejected alternative: deviate silently or block A2 on a note rewrite. Revert: git revert 5b6cd7687 (bridge + migrations only).
  2. D11 "refuses explicitly". A JAX analysis with number_of_cores>1 runs on one core with exactly one INFO line and search.summary records Number of cores = 1 (requested N): <reason>; the user's value is never rewritten. Chosen over raising so Nautilus_jax/Dynesty_jax scripts with 2 cores keep working. If "refuses" must mean raise, it is one branch in parallel.effective_number_of_cores.
  3. RawSamples/samples_from_raw are provisional in search/fit_context.py; A3's samples/adapter.py replaces them in the stacked PR.
  4. NonLinearSearch uses a small ABCMeta subclass so a class defining run counts as concrete (the registry test needs the base to stay abstract); paths/abstract.py::save_summary gained two output-only lines for the cores record.

API Changes

Added Fitness.objective, Fitness(batched=, compile=), Fitness.call_wrap_value_and_grad, module autofit.non_linear.objective, NonLinearSearch.{make_fitness,start_points,run,raw_samples_from,info_from,fitness_overrides}, parallel.{PoolFactory,effective_number_of_cores,check_factor_search_cores,jax_backend_initialized}, search.fit_context.{FitContext,RawSamples,samples_from_raw,UpdateSchedule}, capabilities.run_summary_from. Changed: every call_wrap path lazily jits a JAX analysis (a use_jax=True analysis with a non-traceable likelihood now fails to trace instead of running eagerly); BFGS/LBFGS on JAX use jac=True; Fitness._jit/_vmap are properties over the cache; Drawer and Nautilus no longer define _fit/samples_via_internal_from. Deprecated: Fitness(use_jax_vmap=, use_jax_jit=), Dynesty*(use_jax_jit=), Nautilus(use_jax_vmap=), SettingsSearch(use_jax_vmap=) (default None; search_dict keys unchanged). Removed: test_quick_update_wiring.py (replaced by test_make_fitness.py). Golden identifiers unchanged.
See full details below.

Test Plan

  • pytest test_autofit -x: 3453 passed, 1 skipped, 9 xfailed (A0's strict xfails unchanged)
  • nojax emulation test_autofit/non_linear: 1178 passed, 133 skipped, 7 xfailed; import autofit loads no optional backend
  • Compile-count probe: one trace and one cache entry per kind; building an objective compiles nothing
  • Emcee + use_jax=True per call: jitted scalar 0.304 ms, Emcee's path 0.259 ms (0.85×), eager 14.3 ms (47×); guard test asserts ≤2×
  • BFGS-JAX: call_wrap_value_and_grad calls == nfev, nfev < 4·njev (no finite differences); numpy path unchanged
  • test_sneaky_map.py, test_fork_context.py, EP no-multiprocessing test; conformance both layers for Drawer and Nautilus via run(ctx)
  • PyAutoGalaxy 1357 passed; PyAutoLens 832 passed / 1 xfailed
  • afT {Emcee,DynestyStatic,Nautilus,Nautilus_jax,Dynesty_jax}.py with number_of_cores=2 (scratch copies; the scripts hardcode 1) pass; the JAX legs log exactly one fork-rule line. afT Zeus with 2 cores fails identically on main ("Number of contractions exceeded maximum limit"); 1 core passes
  • afW searches/{mcmc,nest,mle}.py; autolens_workspace scripts/imaging/modeling.py (Nautilus, batched compile 60.9 s) under test mode
  • CI green on unittest 3.12 / 3.13 / nojax, docs, generated-search-docs
Full API Changes (for automation & release notes)

Removed

  • test_quick_update_wiring.py (test only; make_fitness always forwards the quick-update cadence)

Added

  • Fitness.objective(kind, compile=None); Fitness(batched=, compile=); Fitness.call_wrap_value_and_grad
  • autofit.non_linear.objective: jax_objective, chunk_slices, evaluate_in_chunks, kind constants
  • NonLinearSearch.make_fitness, .start_points, .run, .raw_samples_from, .info_from, .fitness_overrides
  • autofit.non_linear.parallel.PoolFactory, effective_number_of_cores, check_factor_search_cores, jax_backend_initialized
  • autofit.non_linear.search.fit_context: FitContext, RawSamples (provisional), samples_from_raw (provisional), UpdateSchedule
  • autofit.non_linear.search.capabilities.run_summary_from

Migration

  • Before: Fitness(..., use_jax_vmap=True) → After: fitness.objective("batched"); the knob still works with a FutureWarning
  • Before: DynestyStatic(use_jax_jit=True) / Nautilus(use_jax_vmap=True) → After: omit; the backend is chosen from analysis.is_jax and the declared capabilities
  • Before: a JAX analysis with number_of_cores=4 forked N recompiles → After: runs on 1 core, one INFO line, recorded in search.summary
  • A new search implements run(ctx) + raw_samples_from(model, internal) per docs/design/run_ctx.md Revision 1; existing _fit overrides keep working until A5

Validation checklist (--auto run — plan was not pre-approved)

  • Effective level: safe (header: safe, cap: refactor → safe)
  • Plan: on the issue (refactor: objective factory, PoolFactory, JAX fork rule and the run(ctx) bridge (search-extensibility A2) #1676), written at start, unmodified since
  • Gate: tests 3453 pass / 1 skip / 9 xfail + nojax 1178 pass + PyAutoLens 832 / PyAutoGalaxy 1357 · smoke afT ×5 with 2 cores + afW ×3 + autolens modeling.py exit 0 (afT Zeus@2 cores pre-existing failure on main) · review CLEAN (Fable session over the Opus-written branch; witness basis-cited: one Fitness( site besides NSS, probe/timing/conformance tests re-run by the reviewer, run_ctx.md Revision 1 present, import autofit free of optional backends) · Heart STALE release validation incomplete: no rehearsal for current source
  • Human: plan sound in hindsight? Rule on flagged items 1 and 2
  • Human: diff matches plan (no scope creep)?
  • Human: merge, amend, or reject — then log the outcome

Generated by the PyAutoLabs agent workflow.

🤖 Generated with Claude Code

Jammy2211 and others added 6 commits October 8, 2026 16:06
Fitness.objective(kind) for kind in {scalar, batched, value_and_grad,
batched_value_and_grad}: on numpy the scalar call and a Python batch loop
(grad kinds raise); on JAX one lazy jax.jit per kind over the unjitted
call, cached per kind, stripped on pickle, wrapped in log_on_first_compile,
with a compile=False eager escape hatch. call_wrap dispatches to the scalar
or batched objective, so a JAX analysis is jitted on every call_wrap path
(Emcee, Zeus, Drawer, the initializer), and call_wrap_value_and_grad books
a value-and-gradient call the same way. Fitness(use_jax_vmap/use_jax_jit)
become FutureWarning aliases of batched/compile. The chunk-and-pad batching
helper is lifted into autofit.non_linear.objective.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Dynesty and BFGS evaluate the scalar objective, Nautilus the batched one;
Emcee, Zeus, Drawer and the initializer reach the lazily jitted scalar
objective through call_wrap, which removes their eager JAX. BFGS/LBFGS on
JAX minimise with jac=True through call_wrap_value_and_grad: the exact
gradient, no finite differences, history and quick updates still booked.
NUTS's log density is the scalar objective; SMC and MultiStart compose
their coordinate transforms on the shared unjitted objective, SMC's batch
is built by jax_objective and MultiStart's chunked sweep uses the shared
evaluate_in_chunks helper.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
NonLinearSearch.make_fitness(analysis, model, **overrides) is now the one
Fitness construction site besides NSS (phase A3b): fom_is_log_likelihood,
convert_to_chi_squared and the resample sentinel derive from the declared
objective_target and invalid_value, and the quick-update cadence,
background worker and live-visual settings are always forwarded.
test_quick_update_wiring.py is retired: the AST scan guarded per-site
forwarding of iterations_per_quick_update, which make_fitness now does for
every search; test_make_fitness.py tests that and pins the construction
sites. DynestyStatic/DynestyDynamic use_jax_jit, Nautilus use_jax_vmap and
SettingsSearch.use_jax_vmap default to None and warn (FutureWarning) when
passed; SettingsSearch.search_dict keeps its keys.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
autofit.non_linear.parallel.pool holds effective_number_of_cores (the
rule: a JAX analysis never forks; number_of_cores>1 runs on one core with
one INFO line), PoolFactory (built once per fit after the gates; the
search, initializer and Nautilus/Dynesty/Emcee/Zeus pools all read its
effective count, with the same pool classes and arguments as before) and
check_factor_search_cores (the EP refusal, moved verbatim from optimise).
search.summary records 'Number of cores = effective (requested N)'. The
grid-search job pool applies the rule to its analysis, sensitivity mapping
to a JAX backend already initialised in the parent, and AnalysisPool
refuses a JAX analysis. make_pool/make_sneaky_pool delegate to the fit's
factory. test_is_jax's use_jax=True double now traces (it is jitted).

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
NonLinearSearch.start_points(model, fitness, n) wraps
initializer.samples_from_model with the fit's effective core count and
makes the start-point plot once, from the first point. Emcee, Zeus, BFGS,
NUTS and SMC drop their manual plot_start_point calls; Drawer and the
Dynesty live-point initialisation use it with plot=False, as before.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
…rated

NonLinearSearch._fit is the bridge: a search that defines run(ctx) gets a
FitContext (model, paths, objective, fitness, test_mode_level, pool,
start_points, rng, resume, checkpointer, schedule, update) built after the
gates, is driven through run, and the context is closed on every exit;
a search that overrides _fit is unchanged. Drawer and Nautilus implement
run(ctx) + raw_samples_from; RawSamples/samples_from_raw are a provisional
stand-in for A3's samples adapter. docs/design/run_ctx.md gains
Revision 1 recording what the bridge pins (schedule takes the total
budget, synchronous update, resume/checkpointer None until A3, close,
fitness_overrides/samples_cls hooks). PoolFactory tracks the pools it
builds so a failed run terminates them.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

decision-taken pending-release PR queued for the next release build

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant