Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
20 commits
Select commit Hold shift + click to select a range
dbf192a
Port find_MAP from pymc-extras and align it with pm.sample
velochy Sep 12, 2026
c949146
Register test_scipy_interface.py in the CI test matrix
velochy Sep 12, 2026
177cee4
Store L-BFGS-B's inverse Hessian as its correction pairs instead of d…
velochy Sep 12, 2026
22f3d85
Defer better_optimize and scipy.optimize imports to first use
velochy Sep 29, 2026
c2eb303
Cover find_MAP error branches in tests
velochy Sep 29, 2026
8077e9f
Address review: exact inverse Hessian, robust optimizer_result, powel…
velochy Sep 29, 2026
53a8822
Deprecate find_MAP's dict return; route internal callers through _fin…
velochy Sep 29, 2026
e12b90f
Rename init="map" to init="jitter+map"; init="map" warns
velochy Sep 29, 2026
3f1333a
Keep init="map" unjittered; jitter+map is the jittered variant
velochy Sep 29, 2026
bbd4669
Address second review round
velochy Oct 1, 2026
7574860
Self-review fixes
velochy Oct 1, 2026
9c586e1
Skip a redundant compile in find_MAP's jitter check; fast compile mod…
velochy Oct 1, 2026
8719f71
Use an already-frozen model as is in find_MAP
velochy Oct 1, 2026
ee8a8c6
Drop the freeze_model and gradient_backend options
velochy Oct 1, 2026
7d9d7ad
Pin better-optimize below 0.5
velochy Oct 1, 2026
8ae3c07
Align find_MAP further with pm.sample
velochy Oct 1, 2026
1a21802
Check jittered starts with the optimizer's own compiled loss
velochy Oct 6, 2026
80e3270
Move find_MAP out of tuning/starting.py into pymc/optimization
velochy Oct 6, 2026
e0d22d5
Keep find_MAP's dict return and unjittered start as defaults, with a …
velochy Oct 7, 2026
8f87b56
Move find_MAP back to tuning/starting.py
velochy Oct 7, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .github/workflows/tests.yml
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,7 @@ jobs:
- |
tests/tuning/test_scaling.py
tests/tuning/test_starting.py
tests/tuning/test_scipy_interface.py
tests/distributions/test_dist_math.py
tests/distributions/test_transform.py
tests/sampling/test_mcmc.py
Expand Down
4 changes: 4 additions & 0 deletions ARCHITECTURE.md
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,10 @@ of the topics below refer to that specific library
* Model comparison, particularly efficient leave-one-out cross-validation approximation
* Data structures for Bayesian inference data storage and manipulation

### SciPy and better-optimize
* Numerical optimization for {func}`pymc.find_MAP`, through `scipy.optimize` wrapped by
`better-optimize` (fused objectives, progress bars, early stopping)


# Modules
The codebase of PyMC is split among single Python file modules at the root
Expand Down
1 change: 1 addition & 0 deletions conda-envs/environment-alternative-backends.yml
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ channels:
dependencies:
# Base dependencies
- arviz>=1.1.0,<2.0
- better-optimize>=0.4.2,<0.5
- blas
- cachetools>=4.2.1,<7
- cloudpickle
Expand Down
1 change: 1 addition & 0 deletions conda-envs/environment-dev.yml
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ channels:
dependencies:
# Base dependencies
- arviz>=1.1.0,<2.0
- better-optimize>=0.4.2,<0.5
- blas
- cachetools>=4.2.1,<7
- cloudpickle
Expand Down
1 change: 1 addition & 0 deletions conda-envs/environment-docs.yml
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ channels:
dependencies:
# Base dependencies
- arviz>=1.1.0,<2.0
- better-optimize>=0.4.2,<0.5
- cachetools>=4.2.1,<7
- cloudpickle
- numpy>=1.25.0
Expand Down
1 change: 1 addition & 0 deletions conda-envs/environment-test.yml
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ channels:
dependencies:
# Base dependencies
- arviz>=1.1.0,<2.0
- better-optimize>=0.4.2,<0.5
- blas
- cachetools>=4.2.1,<7
- cloudpickle
Expand Down
1 change: 1 addition & 0 deletions conda-envs/windows-environment-dev.yml
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ channels:
dependencies:
# Base dependencies (see install guide for Windows)
- arviz>=1.1.0,<2.0
- better-optimize>=0.4.2,<0.5
- blas
- cachetools>=4.2.1,<7
- cloudpickle
Expand Down
1 change: 1 addition & 0 deletions conda-envs/windows-environment-test.yml
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ channels:
dependencies:
# Base dependencies (see install guide for Windows)
- arviz>=1.1.0,<2.0
- better-optimize>=0.4.2,<0.5
- blas
- cachetools>=4.2.1,<7
- cloudpickle
Expand Down
4 changes: 2 additions & 2 deletions pymc/gp/gp.py
Original file line number Diff line number Diff line change
Expand Up @@ -610,7 +610,7 @@ def predict(
R"""
Return mean and covariance of the conditional distribution given a `point`.

The `point` might be the MAP estimate or a sample from a trace.
The `point` is a ``{name: value}`` dict, e.g. one draw of ``pm.find_MAP()`` or ``pm.sample()``.

Parameters
----------
Expand Down Expand Up @@ -1270,7 +1270,7 @@ def predict(self, Xnew, point=None, diag=False, pred_noise=False, model=None):
R"""
Return mean and covariance of the conditional distribution given a `point`.

The `point` might be the MAP estimate or a sample from a trace.
The `point` is a ``{name: value}`` dict, e.g. one draw of ``pm.find_MAP()`` or ``pm.sample()``.

Parameters
----------
Expand Down
12 changes: 6 additions & 6 deletions pymc/model/transform/conditioning.py
Original file line number Diff line number Diff line change
Expand Up @@ -262,13 +262,13 @@ def change_value_transforms(
w = pm.Binomial("w", n=9, p=p, observed=6)

with change_value_transforms(base_m, {"p": logodds}) as transformed_p:
mean_q = pm.find_MAP()
mean_q = pm.find_MAP(return_inferencedata=True, jitter=True).posterior["p"].item()

with change_value_transforms(transformed_p, {"p": None}) as untransformed_p:
new_p = untransformed_p["p"]
std_q = ((1 / pm.find_hessian(mean_q, vars=[new_p])) ** 0.5)[0]
std_q = ((1 / pm.find_hessian({"p": mean_q}, vars=[new_p])) ** 0.5)[0]

print(f" Mean, Standard deviation\\np {mean_q['p']:.2}, {std_q[0]:.2}")
print(f" Mean, Standard deviation\\np {mean_q:.2}, {std_q[0]:.2}")
# Mean, Standard deviation
# p 0.67, 0.16

Expand Down Expand Up @@ -343,12 +343,12 @@ def remove_value_transforms(
with pm.Model() as transformed_m:
p = pm.Uniform("p", 0, 1)
w = pm.Binomial("w", n=9, p=p, observed=6)
mean_q = pm.find_MAP()
mean_q = pm.find_MAP(return_inferencedata=True, jitter=True).posterior["p"].item()

with remove_value_transforms(transformed_m) as untransformed_m:
new_p = untransformed_m["p"]
std_q = ((1 / pm.find_hessian(mean_q, vars=[new_p])) ** 0.5)[0]
print(f" Mean, Standard deviation\\np {mean_q['p']:.2}, {std_q[0]:.2}")
std_q = ((1 / pm.find_hessian({"p": mean_q}, vars=[new_p])) ** 0.5)[0]
print(f" Mean, Standard deviation\\np {mean_q:.2}, {std_q[0]:.2}")

# Mean, Standard deviation
# p 0.67, 0.16
Expand Down
2 changes: 1 addition & 1 deletion pymc/sampling/forward.py
Original file line number Diff line number Diff line change
Expand Up @@ -672,7 +672,7 @@ def sample_posterior_predictive(
Parameters
----------
trace : backend, list, Dataset, DataTree, or MultiTrace
Trace generated from MCMC sampling, or a list of dicts (eg. points or from :func:`~pymc.find_MAP`),
Trace generated from MCMC sampling or :func:`~pymc.find_MAP`, or a list of dicts (eg. points),
or :class:`xarray.Dataset` (eg. DataTree.posterior or DataTree.prior)
model : BaseModel (optional if in ``with`` context)
Model to be used to generate the posterior predictive samples. It will
Expand Down
52 changes: 37 additions & 15 deletions pymc/sampling/mcmc.py
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,7 @@
from pymc.step_methods.arraystep import BlockedStep, PopulationArrayStepShared
from pymc.step_methods.compound import flatten_steps
from pymc.step_methods.hmc import quadpotential
from pymc.tuning.starting import _find_MAP_point
from pymc.util import (
RandomSeed,
RandomState,
Expand Down Expand Up @@ -722,8 +723,8 @@ def sample(
Only applicable to the pymc nuts sampler.
jitter_max_retries : int
Maximum number of repeated attempts (per chain) at creating an initial matrix with uniform
jitter that yields a finite probability. This applies to ``jitter+adapt_diag`` and
``jitter+adapt_full`` init methods.
jitter that yields a finite probability. This applies to ``jitter+adapt_diag``,
``jitter+adapt_full`` and ``jitter+map`` init methods.
n_init : int
Number of iterations of initializer. Only works for 'ADVI' init methods.
trace : backend, optional
Expand Down Expand Up @@ -1700,7 +1701,7 @@ def _init_jitter(
jitter_max_retries: int,
logp_fn: Callable[[PointType], np.ndarray] | None = None,
) -> list[PointType]:
"""Apply a uniform jitter in [-1, 1] to the test value as starting point in each chain.
"""Apply a uniform jitter in [-1, 1] to the initial point as starting point in each chain.

``model.check_start_vals`` is used to test whether the jittered starting
values produce a finite log probability. Invalid values are resampled
Expand Down Expand Up @@ -1786,22 +1787,25 @@ def init_nuts(
Currently, this is ``jitter+adapt_diag``, but this can change in the future. If you
depend on the exact behaviour, choose an initialization method explicitly.
* adapt_diag: Start with a identity mass matrix and then adapt a diagonal based on the
variance of the tuning samples. All chains use the test value (usually the prior mean)
variance of the tuning samples. All chains use the initial point (usually the prior mean)
as starting point.
* jitter+adapt_diag: Same as ``adapt_diag``, but use test value plus a uniform jitter in
[-1, 1] as starting point in each chain.
* jitter+adapt_diag: Same as ``adapt_diag``, but use the initial point plus a uniform
jitter in [-1, 1] as starting point in each chain.
* jitter+adapt_diag_grad:
An experimental initialization method that uses information from gradients and samples
during tuning.
* advi+adapt_diag: Run ADVI and then adapt the resulting diagonal mass matrix based on the
sample variance of the tuning samples.
* 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 model's initial point, as starting point.
This is discouraged.
* jitter+map: Same as ``map``, but search from the initial point plus a uniform jitter
in [-1, 1].
* adapt_full: Adapt a dense mass matrix using the sample covariances. All chains use the
test value (usually the prior mean) as starting point.
* jitter+adapt_full: Same as ``adapt_full``, but use test value plus a uniform jitter in
[-1, 1] as starting point in each chain.
initial point (usually the prior mean) as starting point.
* jitter+adapt_full: Same as ``adapt_full``, but use the initial point plus a uniform
jitter in [-1, 1] as starting point in each chain.

chains : int
Number of jobs to start.
Expand All @@ -1817,8 +1821,8 @@ def init_nuts(
Whether or not to display a progressbar for advi sampling.
jitter_max_retries : int
Maximum number of repeated attempts (per chain) at creating an initial matrix with uniform jitter
that yields a finite probability. This applies to ``jitter+adapt_diag`` and ``jitter+adapt_full``
init methods.
that yields a finite probability. This applies to ``jitter+adapt_diag``, ``jitter+adapt_full``
and ``jitter+map`` init methods.
**kwargs : keyword arguments
Extra keyword arguments are forwarded to pymc.NUTS.

Expand Down Expand Up @@ -1879,6 +1883,8 @@ def model_logp_fn(ip: PointType) -> np.ndarray:
)

apoints = [DictToArrayBijection.map(point) for point in initial_points]
# MAP-based inits run a single search, from the first chain's initvals if given per chain
map_initvals = initvals if initvals is None or isinstance(initvals, dict) else initvals[0]
apoints_data = [apoint.data for apoint in apoints]
potential: quadpotential.QuadPotential

Expand Down Expand Up @@ -1958,7 +1964,15 @@ def model_logp_fn(ip: PointType) -> np.ndarray:
cov = approx.std.eval() ** 2
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(

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

jitter=False. advi_map was unjittered before, and only jitter+map should jitter.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done in 38a4586.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

same initvals forwarding as mcmc.py:1995

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Same commit (38a4586).

model=model,
initvals=map_initvals,
jitter=False,
jitter_max_retries=jitter_max_retries,
random_seed=random_seed_list[0],
progressbar=progressbar and not quiet,
compile_kwargs=compile_kwargs,
)
approx = pm.MeanField(model=model, start=start)
pm.fit(
random_seed=random_seed_list[0],
Expand All @@ -1978,8 +1992,16 @@ def model_logp_fn(ip: PointType) -> np.ndarray:
]
cov = approx.std.eval() ** 2
potential = quadpotential.QuadPotentialDiag(cov, rng=random_seed_list[0])
elif init == "map":
start = pm.find_MAP(include_transformed=True, seed=random_seed_list[0])
elif init in ("map", "jitter+map"):
start = _find_MAP_point(

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The initvals given to pm.sample never reach the MAP search. Forward them.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done in 38a4586: map, jitter+map and advi_map get pm.sample's initvals (chain 0's when given per chain, since there is a single MAP search); test added.

model=model,
initvals=map_initvals,
jitter=init == "jitter+map",
jitter_max_retries=jitter_max_retries,
random_seed=random_seed_list[0],
progressbar=progressbar and not quiet,
compile_kwargs=compile_kwargs,
)
cov = -pm.find_hessian(point=start, negate_output=False)
initial_points = [start] * chains
potential = quadpotential.QuadPotentialFull(cov, rng=random_seed_list[0])
Expand Down
Loading
Loading