Repository navigation
refactor: Fitness.objective factory, PoolFactory + JAX fork rule, start_points and the run(ctx) bridge (search-extensibility A2) - #1679
Merged
Conversation
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>
This was referenced Oct 8, 2026
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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=Falseescape 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 passjac=and stop finite-differencing; Nautilus and SMC take the batched objective; NUTS the scalar log-density.make_fitness(analysis, model, **overrides)builds from the declaredobjective_target/invalid_value;use_jax_jit/use_jax_vmapbecome deprecated aliases (FutureWarning when passed; defaults nowNone).parallel.PoolFactoryowns 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 minimalFitContext/run(ctx)bridge ships with Drawer and Nautilus migrated as proofs (every other_fitkeeps working). Six commits in the plan's order.--autorun at effective level safe; tier judge → human/prm. A3 (#1677) stacks on this branch; A3b on both.Flagged items for the reviewer
docs/design/run_ctx.mdwas 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.updatehands the live object, synchronously;ctx.resume/ctx.checkpointerwired by A3;ctx.rngis aSeedSequenceof the search seed;ctx.close(failed)added. Plus search-side hooksfitness_overrides(analysis)(Nautilus'sbatched/batch_size) andsamples_cls. Rejected alternative: deviate silently or block A2 on a note rewrite. Revert:git revert 5b6cd7687(bridge + migrations only).number_of_cores>1runs on one core with exactly one INFO line andsearch.summaryrecordsNumber 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 inparallel.effective_number_of_cores.RawSamples/samples_from_raware provisional insearch/fit_context.py; A3'ssamples/adapter.pyreplaces them in the stacked PR.NonLinearSearchuses a smallABCMetasubclass so a class definingruncounts as concrete (the registry test needs the base to stay abstract);paths/abstract.py::save_summarygained two output-only lines for the cores record.API Changes
Added
Fitness.objective,Fitness(batched=, compile=),Fitness.call_wrap_value_and_grad, moduleautofit.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: everycall_wrappath lazily jits a JAX analysis (ause_jax=Trueanalysis with a non-traceable likelihood now fails to trace instead of running eagerly); BFGS/LBFGS on JAX usejac=True;Fitness._jit/_vmapare 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=)(defaultNone;search_dictkeys unchanged). Removed:test_quick_update_wiring.py(replaced bytest_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)test_autofit/non_linear: 1178 passed, 133 skipped, 7 xfailed;import autofitloads no optional backenduse_jax=Trueper call: jitted scalar 0.304 ms, Emcee's path 0.259 ms (0.85×), eager 14.3 ms (47×); guard test asserts ≤2×call_wrap_value_and_gradcalls ==nfev,nfev < 4·njev(no finite differences); numpy path unchangedtest_sneaky_map.py,test_fork_context.py, EP no-multiprocessing test; conformance both layers for Drawer and Nautilus viarun(ctx){Emcee,DynestyStatic,Nautilus,Nautilus_jax,Dynesty_jax}.pywithnumber_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 passessearches/{mcmc,nest,mle}.py; autolens_workspacescripts/imaging/modeling.py(Nautilus, batched compile 60.9 s) under test modeFull API Changes (for automation & release notes)
Removed
test_quick_update_wiring.py(test only;make_fitnessalways forwards the quick-update cadence)Added
Fitness.objective(kind, compile=None);Fitness(batched=, compile=);Fitness.call_wrap_value_and_gradautofit.non_linear.objective:jax_objective,chunk_slices,evaluate_in_chunks, kind constantsNonLinearSearch.make_fitness,.start_points,.run,.raw_samples_from,.info_from,.fitness_overridesautofit.non_linear.parallel.PoolFactory,effective_number_of_cores,check_factor_search_cores,jax_backend_initializedautofit.non_linear.search.fit_context:FitContext,RawSamples(provisional),samples_from_raw(provisional),UpdateScheduleautofit.non_linear.search.capabilities.run_summary_fromMigration
Fitness(..., use_jax_vmap=True)→ After:fitness.objective("batched"); the knob still works with a FutureWarningDynestyStatic(use_jax_jit=True)/Nautilus(use_jax_vmap=True)→ After: omit; the backend is chosen fromanalysis.is_jaxand the declared capabilitiesnumber_of_cores=4forked N recompiles → After: runs on 1 core, one INFO line, recorded insearch.summaryrun(ctx)+raw_samples_from(model, internal)perdocs/design/run_ctx.mdRevision 1; existing_fitoverrides keep working until A5Validation checklist (--auto run — plan was not pre-approved)
Fitness(site besides NSS, probe/timing/conformance tests re-run by the reviewer,run_ctx.mdRevision 1 present,import autofitfree of optional backends) · Heart STALErelease validation incomplete: no rehearsal for current sourceGenerated by the PyAutoLabs agent workflow.
🤖 Generated with Claude Code