Repository navigation
Modernize find_MAP - #8429
Modernize find_MAP#8429
Conversation
|
CI note: the two failing test jobs ( |
|
My AI friend doing most of the work, as always - but I am very much here babysitting it. So far this arrangement has been ok so I hope it remains so :) |
Documentation build overview
32 files changed ·
|
1b80464 to
d3ff5af
Compare
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## main #8429 +/- ##
==========================================
+ Coverage 91.95% 92.06% +0.11%
==========================================
Files 128 129 +1
Lines 21276 21453 +177
==========================================
+ Hits 19564 19751 +187
+ Misses 1712 1702 -10
🚀 New features to boost your workflow:
|
|
@jessegrabowski have you had a chance to take a look? |
jessegrabowski
left a comment
There was a problem hiding this comment.
MAP points match pymc_extras.find_MAP to 1e-10 for L-BFGS-B, BFGS, Newton-CG, trust-ncg and powell, and the tuning, init="map" and GP suites pass locally. The blockers are the optimizer_result crashes for nelder-mead and trust-constr.
| if method == "L-BFGS-B" and optimizer_result is not None: | ||
| if hess_inv is None: | ||
| return None | ||
| return np.stack([hess_inv(basis[:, i]) for i in range(n_vars)], axis=-1) |
There was a problem hiding this comment.
not a blocker, existing bug from pymc-extras _compute_inverse_hessian. optional to fix here or in a follow-up. This densifies L-BFGS's 10-pair approximation, not the inverse Hessian. On a 6-parameter regression the diagonal is off by ~1400x vs trust-exact. With compute_hessian=True, compile f_hessp and take the exact route, like the powell case already does.
There was a problem hiding this comment.
Done here rather than in a follow-up (074dd79): compute_hessian always inverts the exact Hessian (fused or hessp); the optimizers' approximations only live in optimizer_result.
| hess_inv = getattr(inner_result, "hess_inv", None) | ||
|
|
||
| if method == "BFGS" and optimizer_result is not None: | ||
| return hess_inv |
There was a problem hiding this comment.
not a blocker, same pymc-extras origin as line 277: BFGS hess_inv is also an approximation, ~1% off in the same model
There was a problem hiding this comment.
Same fix as the L-BFGS-B one (074dd79): BFGS's hess_inv is no longer used for fit.covariance_matrix.
| return None | ||
| return np.stack([hess_inv(basis[:, i]) for i in range(n_vars)], axis=-1) | ||
| if f_hessp is not None: | ||
| H = np.stack([f_hessp(x_star, basis[:, i]) for i in range(n_vars)], axis=-1) |
There was a problem hiding this comment.
raises TypeError under floatX="float32": x_star and basis are float64 and compile doesn't downcast. cast both to floatX
There was a problem hiding this comment.
Fixed in 074dd79: inputs are cast to floatX, with a float32 test.
|
|
||
| if "lowest_optimization_result" in result: | ||
| # basinhopping nests the inner optimizer's result; flatten it over the outer fields | ||
| result = OptimizeResult( |
There was a problem hiding this comment.
merge outer over inner. this replaces basinhopping's totals with one inner run's, nit/nfev go from 5/18 to 2/3 on a 5-hop run
There was a problem hiding this comment.
Fixed in 074dd79: basinhopping's outer totals now win over the inner run's; the test checks nit.
| ) | ||
| vec: tuple[str, ...] = ("variables",) | ||
| mat: tuple[str, ...] = ("variables", "variables_aux") | ||
| data = {"method": xr.DataArray(method)} |
There was a problem hiding this comment.
trust-constr returns its own method key, which overwrites this
There was a problem hiding this comment.
Fixed in 074dd79: our method is written last, so trust-constr's sub-method no longer overwrites it.
| and values[value.name].shape == values[rv.name].shape | ||
| } | ||
| idata_kwargs = {**idata_kwargs, "dims": dims | idata_kwargs.get("dims", {})} | ||
| idata = to_inference_data( |
There was a problem hiding this comment.
nit: this adds an empty sample_stats group
There was a problem hiding this comment.
Fixed in 074dd79: the empty group is dropped.
| potential = quadpotential.QuadPotentialDiag(cov, rng=random_seed_list[0]) | ||
| elif init == "map": | ||
| start = pm.find_MAP(include_transformed=True, seed=random_seed_list[0]) | ||
| start = pm.find_MAP( |
There was a problem hiding this comment.
init="map" used to start from the unjittered initial point. intended?
There was a problem hiding this comment.
I'd say so. 0 can be a saddle point for some common and sensible models (like factor analysis or IRT), so jitter tends to be safer and more robust to beginners in my opinion
There was a problem hiding this comment.
consider init = "jitter+map" for consistency/clarity and warn on init == "map" that this is renamed to jitter+map and use that to silence the warning?
Maybe too much, just a suggestion.
There was a problem hiding this comment.
but for unified syntax, I made it have jitter+map and map separately.
| find_MAP() | ||
| with pytest.warns(UserWarning, match="Gradient not available"): | ||
| find_MAP(model=model, progressbar=False) | ||
| with pytest.raises(Exception): |
There was a problem hiding this comment.
pytest.raises(NotImplementedError)
| ("BFGS", True, False, False), | ||
| ("L-BFGS-B", True, False, False), | ||
| ("trust-exact", True, True, False), | ||
| ("powell", False, False, False), |
There was a problem hiding this comment.
add nelder-mead and trust-constr
There was a problem hiding this comment.
Added both to the parametrized test in 074dd79.
| posterior = idata.posterior.dataset.squeeze(["chain", "draw"]) | ||
| assert posterior["mu"].shape == () and posterior["sigma"].shape == () | ||
| assert ("sigma_log__" in posterior) == include_transformed | ||
| assert ("covariance_matrix" in idata.fit) == compute_hessian |
There was a problem hiding this comment.
not a blocker: check the values against the exact inverse Hessian, not only presence
There was a problem hiding this comment.
Done in 074dd79: the covariance is compared against the inverse of compile_d2logp(jacobian=False) at the optimum.
|
My AI has made all the fixes you requested. Want to do another round @jessegrabowski ? |
jessegrabowski
left a comment
There was a problem hiding this comment.
Want to do another round @jessegrabowski ?
Be careful what you wish for.
| H = np.stack([np.asarray(f_hessp(x_star, e)) for e in basis], axis=-1) | ||
| else: | ||
| raise ValueError("Either `f_hessp` or a fused hessian (`use_hess=True`) is required.") | ||
| return np.linalg.inv(get_nearest_psd(np.asarray(H, dtype="float64"))) |
There was a problem hiding this comment.
Now that this is the exact Hessian it can be indefinite. Clipping negative eigenvalues to 0 then calling inv returns entries around 1e16, or raises LinAlgError after the whole optimization, when the optimum isn't a minimum. Warn and return NaNs (or raise a clear error) when H isn't PD.
There was a problem hiding this comment.
Real jesse here: This bot comment slipped past me. get_nearest_psd is doing the right thing. The robot was upset because it clips negative eigenvalues to exactly zero, which can result in a singular/non-invertible covariance matrix. Some thoughts:
- it's an existing thing so not necessarily your problem.
- It might be right to remove it, but I kind of like it. But I also wrote it so another view would be nice.
- My original logic was that if we have some small numerical noise, it's a shame to throw away the entire (expensive) estimation. Instead we should just give the best possible answer
- If you agree, we should consider clipping the negative eigenvalues to np.spacing(1, dtype=floatX) or something, and add a test (as requested below)
There was a problem hiding this comment.
Agreed, kept the projection (38a4586). Eigenvalues are now clipped to n * eps(floatX) * max|λ| (the np.linalg.matrix_rank tolerance) instead of 0, so noise neither discards the estimate nor makes it singular. A warning fires when the Hessian is clearly indefinite (smallest eigenvalue below -tol), i.e. the point is a saddle rather than a minimum. get_nearest_psd is folded into _compute_inverse_hessian.
|
|
||
|
|
||
| @pytest.mark.parametrize("use_hess", [True, False]) | ||
| def test_compute_inverse_hessian_is_exact(use_hess): |
There was a problem hiding this comment.
Added test_compute_inverse_hessian_indefinite in 38a4586.
| ) | ||
| var_names = [str(var.name) for var in model.value_vars if var.name in names] | ||
| if freeze_model: | ||
| model = freeze_dims_and_data(model) |
There was a problem hiding this comment.
initvals keyed by a TensorVariable are silently ignored once the model is frozen: {x: -50} starts from 0. Convert the keys to names before freezing.
There was a problem hiding this comment.
Fixed in 38a4586: initvals keys are converted to names before freezing; test added.
| potential = quadpotential.QuadPotentialDiag(cov, rng=random_seed_list[0]) | ||
| elif init == "advi_map": | ||
| start = pm.find_MAP(include_transformed=True, seed=random_seed_list[0]) | ||
| start = _find_MAP_point( |
There was a problem hiding this comment.
jitter=False. advi_map was unjittered before, and only jitter+map should jitter.
| monkeypatch.setattr( | ||
| pm.sampling.mcmc, | ||
| "_find_MAP_point", | ||
| lambda **kwargs: jitters.append(kwargs["jitter"]) or find_map_point(**kwargs), |
| else: | ||
| cost_func = CostFuncWrapper(maxeval, progressbar, progressbar_theme, logp_func) | ||
| compute_gradient = False | ||
| res = minimize( |
There was a problem hiding this comment.
Warn when res.success is False. maxiter=1 returns a point far from the optimum with no warning.
There was a problem hiding this comment.
Done in 38a4586: a UserWarning with the optimizer's message.
| f"`{old}` is deprecated, use `{new}` instead.", FutureWarning, stacklevel=2 | ||
| ) | ||
| optimizer_kwargs[new] = optimizer_kwargs.pop(old) | ||
| initvals = optimizer_kwargs.pop("initvals", initvals) |
There was a problem hiding this comment.
start= silently overrides an explicit initvals=. Raise when both are passed.
There was a problem hiding this comment.
For this (and the next comment): The principle should be to harmonize arguments between different fit methods. Take pm.sample as the "truth" and rename stuff? Open to pushback here.
There was a problem hiding this comment.
Agreed that pm.sample is the reference: start became initvals, seed became random_seed, and maxeval maps to scipy's maxiter. Passing an old name together with its replacement now raises (38a4586).
| ) | ||
| optimizer_kwargs[new] = optimizer_kwargs.pop(old) | ||
| initvals = optimizer_kwargs.pop("initvals", initvals) | ||
| random_seed = optimizer_kwargs.pop("random_seed", random_seed) |
There was a problem hiding this comment.
same for seed= / random_seed=, and maxeval= / maxiter=
There was a problem hiding this comment.
Same fix (38a4586): all three pairs raise.
| use_grad = use_grad if use_grad is not None else method_info["uses_grad"] | ||
|
|
||
| if use_hessp is not None and use_hess is None: | ||
| use_hess = not use_hessp |
There was a problem hiding this comment.
not a blocker: use_hessp=False alone turns on the dense Hessian, even for L-BFGS-B. Leave use_hess at the method default instead.
There was a problem hiding this comment.
I think the use_** machinery is a bit scuffed in pymc-extras, it can be rewritten to be less terrible.
There was a problem hiding this comment.
More than the use_hessp=False case: setting one Hessian flag flips the other regardless of method. use_hess=False on L-BFGS-B compiles a hessp it can't use. use_hessp=True on trust-exact drops the Hessian it requires, and find_MAP returns the unoptimized start point without raising. Resolve each None flag from the method independently and respect explicit values:
method_info = MINIMIZE_MODE_KWARGS[method]
hess_was_none, hessp_was_none = use_hess is None, use_hessp is None
use_grad = method_info["uses_grad"] if use_grad is None else use_grad
use_hess = method_info["uses_hess"] if hess_was_none else use_hess
use_hessp = method_info["uses_hessp"] if hessp_was_none else use_hessp
if use_hess and use_hessp:
if hessp_was_none and not hess_was_none:
use_hessp = False
else:
use_hess = FalseWarn when an explicit True names something the method doesn't use.
There was a problem hiding this comment.
Implemented your resolution in 38a4586, with one addition: a flag the method can't use is ignored with a warning, otherwise use_hessp=True on trust-exact would still drop the Hessian it needs. use_grad gets the same treatment.
| include_transformed : bool, default False | ||
| Whether to also return the values of transformed (unconstrained) variables, e.g. ``sigma_log__``. | ||
| compute_hessian : bool, default False | ||
| Compute the inverse hessian at the optimum and store it as ``fit.covariance_matrix``. |
There was a problem hiding this comment.
The covariance is the inverse Hessian of logp(jacobian=False) over the unconstrained value variables, which isn't the Laplace covariance in that space. State the space and drop the Laplace sentence.
There was a problem hiding this comment.
Done in 38a4586: the docstring states the space, and the Laplace sentence is gone.
| # Variable keys would not match the frozen model's variables, so key by name | ||
| initvals = initvals and {getattr(k, "name", k): v for k, v in initvals.items()} | ||
| if freeze_model: | ||
| if freeze_model and not isinstance(model, FrozenModel): |
There was a problem hiding this comment.
I'm not in love with the freeze model option, I'm willing to wager that @ricardoV94 will object on the grounds that if a user wants it, he can just do it.
There was a problem hiding this comment.
Removed in 9206e69. The model is used as given, so passing one from pm.model.transform.freeze_model gets constant folding and cached compiled functions. The JAX backend works on unfrozen models too.
|
|
||
| """Compile model log-densities into the callables expected by ``scipy.optimize``.""" | ||
|
|
||
| from __future__ import annotations |
There was a problem hiding this comment.
is this being used? llms love this import but it doesn't do anything 99% of the time
| model : Model (optional if in ``with`` context) | ||
| backend : str, optional | ||
| Computational backend, one of "numba", "c" or "jax". Defaults to the PyTensor default mode. | ||
| compile_kwargs : dict, optional |
There was a problem hiding this comment.
i recommend just removing the gradient_backend thing. It is a vestigial feature from when our Op coverage was less complete.
There was a problem hiding this comment.
Removed in 9206e69: derivatives always come from PyTensor, and backend="jax" still compiles via JAX. That also let the private compile helper fold into scipy_optimize_funcs_from_loss.
| dependencies: | ||
| # Base dependencies | ||
| - arviz>=1.1.0,<2.0 | ||
| - better-optimize>=0.4.2,<1.0 |
There was a problem hiding this comment.
pin this <0.5 everywhere, 1.0 is too aggressive
There was a problem hiding this comment.
Pinned to >=0.4.2,<0.5 everywhere in a4e415c.
There was a problem hiding this comment.
Small question and suggestion.
Can you give me a one sentence about backward incompatible changes and how users should change their code to include in the release notes?
And clarify if this PR breaks compat immediately or only issues the warning but keeps working.
| * advi: Run ADVI to estimate posterior mean and diagonal mass matrix. | ||
| * advi_map: Initialize ADVI with MAP and use MAP as starting point. | ||
| * map: Use the MAP as starting point. This is discouraged. | ||
| * map: Use the MAP, searched for from the test value, as starting point. This is |
There was a problem hiding this comment.
Wording copied from the neighbouring entries; it now says "the model's initial point" (4a531ee). The older adapt_diag / adapt_full entries still say "test value". Happy to fix those here too if you want.
There was a problem hiding this comment.
Fixed the remaining "test value" entries in init_nuts and _init_jitter (8af31ef).
| outputs.append(grad) | ||
| if use_hess: | ||
| outputs.append(pytensor.gradient.jacobian(grad, [flat_input])[0]) | ||
| f_fused = compile([flat_input], outputs if len(outputs) > 1 else loss, **compile_kwargs) |
There was a problem hiding this comment.
want to do some trust_input=True, somewhere?
There was a problem hiding this comment.
Done in 4a531ee: the loss and hessp run with trust_input=True, behind a cast to the input dtype since scipy hands over float64.
There was a problem hiding this comment.
that defeats part of the trust_input part I guess, but still cheaper? OTOH isn't it always a single input function, why the comprehension?
There was a problem hiding this comment.
Still cheaper: 8.3 µs per call with checked inputs against 5.1 µs trusted with the cast, the same as without it. The comprehension was for hessp's (x, p); the casts are now explicit per function (8af31ef).
| assert_allclose([r["x"], r["y"], r["det"]], [50, 50, 100]) | ||
|
|
||
|
|
||
| def test_find_MAP_frozen_model(normal_model): |
There was a problem hiding this comment.
move these tests (and map functionality) out of tuning/starting.py, unless they are specifically abound starting at the map from pm.sample?
There was a problem hiding this comment.
Moved to pymc/optimization/{map,scipy_interface}.py, with tests in tests/optimization (77acf5c). pymc.find_MAP and pymc.tuning.find_MAP still work. The name is my pick, so let me know if you prefer another home.
Can't this reuse the function that find_MAP used? |
Can we change this / take control over the progressbar? Not a blocker (nor very urgent). @jessegrabowski |
|
@ricardoV94 thanks for the review. Inline comments are answered in their threads; the rest is here. Release notes
Does it break immediately or only warn? Mixed:
"Can't this reuse the function that find_MAP used?" Yes, done in 4a531ee. Jittered starts are checked with the optimizer's own compiled loss, with held-fixed variables as shared inputs so one function serves every redraw. No separate logp is compiled, and Progress bar. Feasible as a follow-up: Three things I'd like your guidance on
|
I have to revisit my categorical reasoning 101, but I feel there ought to be more alternatives than returning the dict for one release alone Yes, I'd say keep the old behavior of dict and no jitter by default with the warning for a few releases. The callable bit seems a bit more niche and we can deprecate immediately. Raise an informative error if one was provided. Also to nit with the bot, adding a dependency is not breaking anything. Must revisit causal analysis 101 as well I guess |
|
Re module location... nhe leave it for now. We should probably get rid of the whole |
|
@velochy didn't hear back from you on suggestion to make transition gradual and not break now |
|
@ricardoV94 done, thanks for the guidance.
CI: the one red job is |
Rebuild pm.find_MAP on the pymc-extras implementation: fused loss/grad/hess functions over one flat parameter vector, better_optimize for the scipy interface, basinhopping and optional inverse-Hessian. Parameters, defaults and conventions follow pm.sample (backend/compile_kwargs, initvals/jitter/ random_seed via _init_jitter, return_inferencedata via to_inference_data). Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
…ensifying it optimizer_result turned scipy's LbfgsInvHessProduct into an n x n matrix via n matvecs: O(n^2) memory (8.7 GB transient at n = 19k) for an object that is 2 m n floats. Keep it as hess_inv_sk / hess_inv_yk with dims (lbfgs_corrections, variables); a generic LinearOperator is still densified as before. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
`import pymc` eagerly imports pymc.tuning, so module-level better_optimize imports pulled in scipy.optimize/linalg/special/stats/interpolate/spatial and broke test_eager_import_heavy_dependencies. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
…l fallback - compute_hessian always uses the exact Hessian (fused or hessp), never BFGS/L-BFGS-B's approximate hess_inv; floatX-safe under float32 - optimizer_result: split ragged tuples (nelder-mead final_simplex), label "variables" only when sizes match (trust-constr's empty jac), keep our method over trust-constr's, and let basinhopping's outer totals win over the inner run's - only fall back to powell for discrete vars when the method uses gradients - drop the empty sample_stats group - tests: nelder-mead/trust-constr, exact covariance values, float32, NotImplementedError Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
…d_MAP_point return_inferencedata=False keeps working but raises a FutureWarning. The optimization moves into _fit_MAP, and init="map"/"advi_map" read the point from the private _find_MAP_point. Docstrings and tests read the point from the returned InferenceData instead. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
The MAP search starts from a jittered point, so name it like the other jitter+ inits and forward jitter_max_retries. Also pin the discrete find_MAP test's randomness, which the jitter default had made flaky. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
- use_* flags: resolve each None from the method independently, respect explicit values, and warn about and ignore flags the method cannot use - inverse Hessian: clip eigenvalues to a relative tolerance instead of 0 and warn when the Hessian is not positive definite (saddle point) rather than raising or returning ~1e16 - warn when the optimizer does not converge - raise when a deprecated kwarg is passed together with its replacement - key initvals by name so Variable keys survive model freezing - advi_map does not jitter; MAP-based inits receive pm.sample's initvals - compute_hessian docstring states the space the covariance lives in Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
- compute_hessian raises for discrete variables instead of returning garbage rows - post-hoc hessp honours gradient_backend="jax" for gradient-free methods - drop the generic LinearOperator densification (only L-BFGS-B returns one, handled compactly) - error message lists canonical method names; shorter private docstrings - jitter_rvs typed as TensorVariable (no new mypy errors in mcmc.py) - gp.predict / sample_posterior_predictive docstrings no longer imply find_MAP returns a dict - ARCHITECTURE.md lists scipy.optimize / better-optimize - tests: compute_hessian once per Hessian route (fewer compiles), discrete-hessian and jax-hessp tests Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
…e in tests The jitter retry loop compiled a full backend logp just to check a few starts for finiteness; a FAST_COMPILE one suffices. test_starting.py compiles with fast_unstable_sampling_mode like the step-method tests: tests/tuning drops from ~55 s to ~29 s (main: ~32 s). Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Refreezing a FrozenModel turned it into a plain Model and dropped its cached functions. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
A user who wants a frozen model can pass one from pm.model.transform.freeze_model, which is used as is. Derivatives are always taken by PyTensor; the JAX-autodiff path was vestigial. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
- include_transformed goes through idata_kwargs as in pm.sample; the keyword warns and is forwarded - progressbar accepts pm.sample's string options - transformed values no longer get their RV's dims, matching pm.sample - compute_hessian compiles hessp alongside the loss instead of recompiling the loss - _fit_MAP has no defaults, so they live only on find_MAP; _find_MAP_point states its settings Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
find_MAP draws its start from the initial point function pm.sample uses and tests it against the loss it compiles anyway, with held-fixed variables as shared inputs. That drops the extra logp compile and the _init_jitter change in mcmc.py. The compiled functions use trust_input. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
MAP estimation is not about NUTS tuning. find_MAP and its scipy interface now live in pymc/optimization, with tests in tests/optimization. pymc.find_MAP is unchanged and pymc.tuning.find_MAP stays importable. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
…FutureWarning Both stay for a few releases; return_inferencedata=True and jitter=True opt in to the future defaults. A callable method raises an informative TypeError. The trust_input casts are explicit per function, and the remaining "test value" wording in init_nuts is fixed. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
The module's new home is to be decided when pymc.tuning is removed. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Port
find_MAPfrom pymc-extras into PyMCCloses #7308.
Summary
pm.find_MAPis rewritten on top of thepymc_extras.find_MAPimplementation and aligned withpm.samplein parameter names, defaults and conventions. The pymc-extras design (fusedloss/gradient/hessian functions compiled over one flat parameter vector,
better_optimizeforthe scipy interface and progress bar,
basinhoppingsupport, optional inverse-Hessian) is takenas the base; where pymc-extras and
pm.sampledisagree,pm.samplewins.Highlights:
pm.sample.backend=/compile_kwargs=go throughresolve_backend_compile_kwargs, so the default backend is whatever PyTensor's default mode is(numba), exactly like NUTS. Derivatives are always taken by PyTensor; pymc-extras'
gradient_backend="jax"option (JAX autodiff on the compiled loss) is dropped as vestigial.pm.sample.initvals=,jitter=,jitter_max_retries=and
random_seed=use the initial point functionpm.sampleuses, with the same redraw rulefor non-finite jittered starts, which also fixes pymc-extras#687 (jitter never applied /
broken by model freezing).
pm.sample.return_inferencedata=Truereturns aDataTreebuilt by
pm.to_inference_datafrom a one-draw trace, soposterior,observed_data,constant_data, deterministics, coords/dims andidata_kwargs(log_likelihood, ...) allbehave like a sampling run. Two extra groups carry the optimizer output:
fit(
mean_vector, optionalcovariance_matrix) andoptimizer_result(every field of the scipyOptimizeResult, per-parameter fields labelled with coordinate-aware names such asbeta[Intercept]). The classic{name: value}dict stays the default for a few releases,with a
FutureWarning."bfgs","Powell"keep working),use_grad/use_hess/use_hesspare user-controllable with method-based defaults, andmethod="basinhopping"is available.Backward compatibility
Existing
find_MAPcalls keep working. For release notes:FutureWarning) but keeps working: the dict return (default, orreturn_inferencedata=False), the unjittered start whenjitteris not given,start→initvals,seed→random_seed,maxeval→maxiter,include_transformed=→idata_kwargs,return_raw, a single positional start dict, andprogressbar_theme(ignored).methodraises aTypeError, and everything exceptmethodis keyword-only, sofind_MAP(start, vars)raises aTypeError.pm.find_MAPandpymc.tuning.find_MAPimport paths, andpm.sample(init="map")/"advi_map", apart from now honouringinitvals.API
start(positional)initvalsinitvalsstart=and a positional dict still work with aFutureWarningseedrandom_seedrandom_seedseed=warnsvarsvarsjitter_rvsjitter+jitter_max_retriespm.samplesemantics; only optimized variables are jittered. Off unless requested for now, with aFutureWarningwhen not givenreturn_rawidata.optimizer_resultmaxevalmaxiter(via kwargs)maxiter(via kwargs)maxeval=warns and is forwardedprogressbar_themebetter_optimize)include_transformed=Trueinclude_transformed=Trueidata_kwargs={"include_transformed": ...}, default Falsepm.sample; the old keyword warns and is forwardedgradient_backendcompile_kwargsbackend+compile_kwargspm.sampleconventionreturn_inferencedata=Truecompute_hessianfreeze_modelpm.model.transform.freeze_modelinsteadBehavioural differences vs. pymc-extras worth knowing:
posteriorwithidata_kwargs={"include_transformed": True},exactly as in
pm.sample, instead of a separateunconstrained_posteriorgroup.progressbaracceptspm.sample's string options, which simply enable the optimizer's bar.constant_datais only present when the model has constant data (converter behaviour).fit/optimizer_resultfollows model definition order.compute_hessian=Trueinverts the exact Hessian of-logp(jacobian=False)over the optimizedunconstrained value variables, from the fused Hessian or from Hessian-vector products compiled
after optimization if the optimizer did not need them. pymc-extras reused BFGS / L-BFGS-B's
hess_inv, which is only an approximation (L-BFGS-B's diagonal can be off by orders ofmagnitude); those approximations are still stored in
optimizer_result. Eigenvalues areclipped to a relative tolerance (
n · eps · max|λ|, as innp.linalg.matrix_rank) insteadof 0, so numerical noise neither discards the estimate nor makes it singular, and a warning
fires when the Hessian is clearly indefinite (a saddle point, not a minimum). With discrete
variables in
varsit raises, since the Hessian is undefined there."powell"with a warning when discrete variables areoptimized or the model has no gradient, only if
use_gradwas left atNone; explicitlyasking for gradients raises. Gradient-free methods such as
nelder-meadare kept as given.use_grad/use_hess/use_hessp: eachNoneis resolved from what the method uses,explicit values are respected, and a flag the method cannot use is ignored with a warning
(so
use_hessp=Trueon trust-exact no longer drops the Hessian it needs). Ofhessandhessponly one is used: the explicit one, elsehessp.UserWarningfires when the optimizer reports that it did not converge.start+initvals,seed+random_seed,maxeval+maxiter) raises instead of silently picking one.initvalsmay be keyed by variables as well as names.extra is compiled for it. Variables held fixed via
varsare shared inputs of that loss.trust_input=True, after a dtype cast.OptimizeResultfield is stored, including method-specific ones (nelder-mead'sfinal_simplexis split into two variables, trust-constr's constraint fields keep their owndims). Basinhopping reports its own
nit/nfevtotals, with the best inner run's fieldsunderneath.
ValueErrorinstead of being forwarded to scipy, and a callablemethodraises aTypeError: it is no longer supported.Large models and the Hessian
MAP estimation is mostly worth reaching for on models too large for full MCMC, so parameter
counts in the tens of thousands are normal here and anything O(n²) is a liability.
find_MAPtherefore never materializes a dense n×n matrix unless asked: L-BFGS-B's inverse Hessian is
stored in
optimizer_resultas its m correction pairs (hess_inv_sk/hess_inv_yk, dims(lbfgs_corrections, variables)) instead of being densified through n matvecs, which costs8.7 GB transient at n = 19k for an object that is only 2·m·n floats. A dense
fit.covariance_matrixis only computed withcompute_hessian=True, anduse_hessdefaults toFalse for every method that can use
hesspinstead.Files
pymc/tuning/scipy_interface.py(new): compile helpers ported from pymc-extras(
scipy_optimize_funcs_from_loss,set_optimizer_function_defaults,_compute_inverse_hessian).get_nearest_psdis folded into_compute_inverse_hessian.pymc/tuning/starting.py: the rewrite plus the idata helpers.better_optimizeandscipy.optimizeare imported on first use soimport pymcstays light.pymc/sampling/mcmc.py:init="map"/"jitter+map"/"advi_map"get the point from theprivate
_find_MAP_point, forwardingcompile_kwargs,progressbarandpm.sample'sinitvals(previously ignored; with per-chaininitvalsthe single MAP search uses chain 0's).mapandadvi_mapkeep their old unjittered start.pymc/model/transform/conditioning.py,pymc/gp/gp.py,pymc/sampling/forward.py:docstrings no longer present
find_MAP's output as a point dict.ARCHITECTURE.md: listsscipy.optimize/better-optimizeunder functionality not in PyMC..github/workflows/tests.yml: registerstests/tuning/test_scipy_interface.py.pyproject.toml,requirements-dev.txt,conda-envs/*.yml: new dependencybetter-optimize>=0.4.2,<0.5(conda-forge and PyPI; depends only on numpy/scipy/rich/pandas/joblib/threadpoolctl, all
already required or trivially available).
tests/tuning/test_starting.py(existing tests adapted, pymc-extras tests ported, newtests for jitter/seeding, legacy kwargs,
idata_kwargs,varssubsets, invalid start),tests/tuning/test_scipy_interface.py(new). Other callers intests/gp,tests/distributionsand
tests/model/transformread the point from the returnedInferenceData.Decisions for reviewers
better-optimizebecomes a hard dependency. It is maintained by @jessegrabowski and iswhat pymc-extras builds on (progress bar, fused-function detection,
maxiter/ tolerancedefaults, early stopping,
basinhopping). Vendoring the partsfind_MAPneeds would beroughly 150 lines of duplicated code, so the dependency looked like the more maintainable
option. It is only imported on first use. If a new dependency is not acceptable, the
fallback is to vendor a minimal
minimizewrapper.FutureWarning;return_inferencedata=Trueopts in to theInferenceData. On the dictpath transformed values are included, as before. Nothing inside PyMC uses the dict:
the MAP-based
initoptions get the point from the private_find_MAP_point, and thedocstrings and tests read it from the
InferenceData.include_transformedmoves intoidata_kwargs, default False for theInferenceData, as inpm.sample. Both oldpymc and pymc-extras had it as a keyword defaulting to
True; the keyword still works with aFutureWarning.jitter=Truewill become the default, likepm.sample, so the default initial pointcannot park the optimizer on a saddle point (pymc-extras#687). For a few releases the start
stays unjittered unless
jitter=Trueis passed, with aFutureWarningwhenjitteris notgiven. Results are reproducible with
random_seed.Variables held fixed via
varsare not jittered.init="map"keeps its old unjittered start, andinit="jitter+map"is new. The lattersearches for the MAP from a jittered point and honours
jitter_max_retries, following theadapt_diag/jitter+adapt_diagpattern.freeze_modeloption, unlike pymc-extras and likepm.sample. The model is used asgiven: a model from
pm.model.transform.freeze_modelgets constant folding and compiledfunctions cached across calls. The JAX backend works on unfrozen models too.
progressbar_themeis dropped.better_optimizeowns the progress bar and does not take arich theme.
Follow-ups (separate PRs)
pymc_extras.find_MAPre-exportpm.find_MAP, importscipy_optimize_funcs_from_loss/set_optimizer_function_defaults/_compute_inverse_hessianfrompymc.tuning.scipy_interfaceinfit_laplaceand DADVI, andread
idata.posteriorinstead ofidata.unconstrained_posterior.find_MAPexample notebooks still index the returned dict.Checks run locally
tests/tuning,tests/test_root_namespace.py,tests/sampling/test_mcmc.py,tests/sampling/test_forward.py,tests/model/transform/test_conditioning.py,tests/distributions/test_truncated.pyandtests/gp/test_gp.pyall pass with the numbadefault backend. JAX paths are exercised where
jaxis installed. Both tuning modules have100% line coverage.
test_starting.pycompiles withpymc.testing.fast_unstable_sampling_modelike the step-method tests (the default backend stays covered by the
init, GP and JAXtests), so
tests/tuningtakes about 40 s, against 32 s on main with 13 tests.pre-commit runover the branch is clean;mypyis clean on the tuning modules and adds noerrors to
mcmc.py.🤖 Generated with Claude Code