diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index 40473a014c..df932576ce 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -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 diff --git a/ARCHITECTURE.md b/ARCHITECTURE.md index 37990d1ac8..e5ae16e9a2 100644 --- a/ARCHITECTURE.md +++ b/ARCHITECTURE.md @@ -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 diff --git a/conda-envs/environment-alternative-backends.yml b/conda-envs/environment-alternative-backends.yml index b18cb4648a..75c1be2da1 100644 --- a/conda-envs/environment-alternative-backends.yml +++ b/conda-envs/environment-alternative-backends.yml @@ -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 diff --git a/conda-envs/environment-dev.yml b/conda-envs/environment-dev.yml index 0ea9e622e0..8bb210064d 100644 --- a/conda-envs/environment-dev.yml +++ b/conda-envs/environment-dev.yml @@ -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 diff --git a/conda-envs/environment-docs.yml b/conda-envs/environment-docs.yml index dd04b40239..99da60e748 100644 --- a/conda-envs/environment-docs.yml +++ b/conda-envs/environment-docs.yml @@ -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 diff --git a/conda-envs/environment-test.yml b/conda-envs/environment-test.yml index 42de245414..90bfd095f9 100644 --- a/conda-envs/environment-test.yml +++ b/conda-envs/environment-test.yml @@ -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 diff --git a/conda-envs/windows-environment-dev.yml b/conda-envs/windows-environment-dev.yml index 3093686b3d..d087225c4d 100644 --- a/conda-envs/windows-environment-dev.yml +++ b/conda-envs/windows-environment-dev.yml @@ -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 diff --git a/conda-envs/windows-environment-test.yml b/conda-envs/windows-environment-test.yml index 13e0a97ce5..85c3bd76a8 100644 --- a/conda-envs/windows-environment-test.yml +++ b/conda-envs/windows-environment-test.yml @@ -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 diff --git a/pymc/gp/gp.py b/pymc/gp/gp.py index 1a0081a5c7..d513b64d3c 100644 --- a/pymc/gp/gp.py +++ b/pymc/gp/gp.py @@ -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 ---------- @@ -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 ---------- diff --git a/pymc/model/transform/conditioning.py b/pymc/model/transform/conditioning.py index 046549e6c8..7a327beb4c 100644 --- a/pymc/model/transform/conditioning.py +++ b/pymc/model/transform/conditioning.py @@ -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 @@ -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 diff --git a/pymc/sampling/forward.py b/pymc/sampling/forward.py index 4b52064dc4..aea50d8370 100644 --- a/pymc/sampling/forward.py +++ b/pymc/sampling/forward.py @@ -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 diff --git a/pymc/sampling/mcmc.py b/pymc/sampling/mcmc.py index 7e206589af..9d6cad5ced 100644 --- a/pymc/sampling/mcmc.py +++ b/pymc/sampling/mcmc.py @@ -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, @@ -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 @@ -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 @@ -1786,10 +1787,10 @@ 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. @@ -1797,11 +1798,14 @@ def init_nuts( 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. @@ -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. @@ -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 @@ -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( + 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], @@ -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( + 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]) diff --git a/pymc/tuning/scipy_interface.py b/pymc/tuning/scipy_interface.py new file mode 100644 index 0000000000..33c981c48e --- /dev/null +++ b/pymc/tuning/scipy_interface.py @@ -0,0 +1,170 @@ +# Copyright 2024 - present The PyMC Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Compile model log-densities into the callables expected by ``scipy.optimize``.""" + +import warnings + +from collections.abc import Callable +from typing import cast + +import numpy as np +import pytensor +import pytensor.tensor as pt + +from pytensor.tensor import TensorVariable + +from pymc.pytensorf import compile, floatX, join_nonshared_inputs, rewrite_pregrad + + +def set_optimizer_function_defaults( + method: str, use_grad: bool | None, use_hess: bool | None, use_hessp: bool | None +) -> tuple[bool, bool, bool]: + """Resolve ``None`` flags from what ``method`` uses, ignoring flags it can't use. + + Of a Hessian and a Hessian-vector product only one is used: the explicit one, else ``hessp``. + """ + from better_optimize.constants import MINIMIZE_MODE_KWARGS + + info = MINIMIZE_MODE_KWARGS[method] + for flag, value, key in ( + ("use_grad", use_grad, "uses_grad"), + ("use_hess", use_hess, "uses_hess"), + ("use_hessp", use_hessp, "uses_hessp"), + ): + if value and not info[key]: + warnings.warn( + f"Method {method!r} does not use `{flag}`; it will be ignored.", UserWarning + ) + hess_was_none, hessp_was_none = use_hess is None, use_hessp is None + use_grad = info["uses_grad"] and (info["uses_grad"] if use_grad is None else use_grad) + use_hess = info["uses_hess"] and (info["uses_hess"] if hess_was_none else use_hess) + use_hessp = info["uses_hessp"] and (info["uses_hessp"] if hessp_was_none else use_hessp) + if use_hess and use_hessp: + if not (hess_was_none or hessp_was_none): + warnings.warn( + "Only one of `use_hess` and `use_hessp` is used; using `use_hessp`.", UserWarning + ) + if hessp_was_none and not hess_was_none: + use_hessp = False + else: + use_hess = False + return bool(use_grad), bool(use_hess), bool(use_hessp) + + +def scipy_optimize_funcs_from_loss( + loss: TensorVariable, + inputs: list[TensorVariable], + initial_point_dict: dict[str, np.ndarray] | None = None, + use_grad: bool | None = None, + use_hess: bool | None = None, + use_hessp: bool | None = None, + compile_kwargs: dict | None = None, + inputs_are_flat: bool = False, +) -> tuple[Callable, Callable | None]: + """Compile ``loss`` into scipy-compatible ``(f_fused, f_hessp)`` callables of one flat vector. + + Parameters + ---------- + loss : TensorVariable + Scalar loss to minimize. + inputs : list of TensorVariable + Input variables, joined into a single raveled vector unless ``inputs_are_flat``. + initial_point_dict : dict, optional + Maps input names to values; only used to determine input shapes. + use_grad, use_hess, use_hessp : bool, optional + Which derivatives to compile into the returned functions. + compile_kwargs : dict, optional + Keyword arguments passed on to :func:`pymc.compile`. + inputs_are_flat : bool + Set when ``inputs`` already is a single flat vector. + + Returns + ------- + f_fused : Callable + Returns the loss, or ``(loss, grad)`` / ``(loss, grad, hess)`` when derivatives are requested. + f_hessp : Callable or None + Hessian-vector product function, if requested. + """ + if use_hess and not use_grad: + raise ValueError("Cannot compute hessian without also computing the gradient") + compile_kwargs = {} if compile_kwargs is None else compile_kwargs + if not isinstance(inputs, list): + inputs = [inputs] + if inputs_are_flat: + [flat_input] = inputs + else: + outputs, flat_input = join_nonshared_inputs( + point=initial_point_dict or {}, outputs=[loss], inputs=inputs + ) + loss = cast(TensorVariable, outputs[0]) + loss = rewrite_pregrad(loss) + + # scipy hands over float64, so cast to the input dtype and let PyTensor skip its input checks + dtype = flat_input.dtype + + f_hessp = None + if use_hessp: + p = pt.tensor("p", shape=flat_input.type.shape, dtype=flat_input.dtype) + hessp = pytensor.gradient.hessian_vector_product(loss, [flat_input], p) + fn_hessp = compile([flat_input, p], hessp[0], trust_input=True, **compile_kwargs) + + def f_hessp(x, p): + return fn_hessp(np.asarray(x, dtype=dtype), np.asarray(p, dtype=dtype)) + + outputs = [loss] + if use_grad: + grad = cast(TensorVariable, pytensor.gradient.grad(loss, flat_input)) + outputs.append(grad) + if use_hess: + outputs.append(pytensor.gradient.jacobian(grad, [flat_input])[0]) + fn = compile( + [flat_input], outputs if len(outputs) > 1 else loss, trust_input=True, **compile_kwargs + ) + return (lambda x: fn(np.asarray(x, dtype=dtype))), f_hessp + + +def _compute_inverse_hessian( + optimal_point: np.ndarray, + f_fused: Callable | None = None, + f_hessp: Callable | None = None, + use_hess: bool = False, +) -> np.ndarray: + """Inverse of the exact Hessian at ``optimal_point`` (never an optimizer's approximation). + + Eigenvalues are clipped to a relative tolerance, warning when clearly negative (not a minimum). + """ + x_star = floatX(np.asarray(optimal_point)) + if use_hess and f_fused is not None: + _, _, H = f_fused(x_star) + elif f_hessp is not None: + basis = floatX(np.eye(len(x_star))) + 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.") + H = np.asarray(H, dtype="float64") + eigval, eigvec = np.linalg.eigh((H + H.T) / 2) + # Same relative tolerance as np.linalg.matrix_rank, at the precision H was computed in + tol = max( + np.abs(eigval).max() * len(eigval) * np.finfo(pytensor.config.floatX).eps, + np.finfo("float64").tiny, + ) + if eigval.min() < -tol: + warnings.warn( + f"The Hessian at the optimum is not positive definite (smallest eigenvalue {eigval.min():.3g}), " + "so the point may be a saddle point rather than a minimum. Its eigenvalues were clipped to " + "compute `fit.covariance_matrix`.", + UserWarning, + ) + return (eigvec / np.maximum(eigval, tol)) @ eigvec.T diff --git a/pymc/tuning/starting.py b/pymc/tuning/starting.py index d9a2bd11d5..ff2269cd36 100644 --- a/pymc/tuning/starting.py +++ b/pymc/tuning/starting.py @@ -12,258 +12,558 @@ # See the License for the specific language governing permissions and # limitations under the License. -""" -Created on Mar 12, 2011. +"""Maximum a posteriori (MAP) estimation.""" -@author: johnsalvatier -""" +from __future__ import annotations import warnings from collections.abc import Sequence +from itertools import product +from typing import TYPE_CHECKING, Any, Literal, NamedTuple, cast import numpy as np +import pytensor import pytensor.gradient as tg +import xarray as xr -from numpy import isfinite +from pytensor.compile import Function from pytensor.graph.basic import Variable -from pytensor.utils import lazy_scipy_module -from rich.console import Console -from rich.progress import Progress, TextColumn - -import pymc as pm - -from pymc.blocking import DictToArrayBijection, RaveledVars -from pymc.initial_point import make_initial_point_fn -from pymc.model import modelcontext -from pymc.progress_bar import CustomProgress, default_progress_theme -from pymc.pytensorf import floatX, inputvars +from pytensor.graph.replace import graph_replace +from pytensor.tensor import TensorVariable +from xarray import DataTree + +from pymc.backends.arviz import to_inference_data +from pymc.backends.base import MultiTrace +from pymc.backends.ndarray import NDArray +from pymc.blocking import DictToArrayBijection, PointType, RaveledVars +from pymc.initial_point import StartDict, make_initial_point_fns_per_chain +from pymc.model import Model, modelcontext +from pymc.progress_bar import ProgressBarOptions +from pymc.pytensorf import inputvars, resolve_backend_compile_kwargs +from pymc.tuning.scipy_interface import ( + _compute_inverse_hessian, + scipy_optimize_funcs_from_loss, + set_optimizer_function_defaults, +) from pymc.util import ( + RandomState, get_default_varnames, + get_random_generator, get_value_vars_from_user_vars, ) from pymc.vartypes import discrete_types, typefilter -optimize = lazy_scipy_module("optimize") +if TYPE_CHECKING: # better_optimize and scipy.optimize are heavy; import on first use + from better_optimize.constants import minimize_method + from scipy.optimize import OptimizeResult __all__ = ["find_MAP"] +_LEGACY_KWARGS = {"start": "initvals", "seed": "random_seed", "maxeval": "maxiter"} + + +def _canonical_method(method: str) -> str: + if callable(method): + raise TypeError( + "A callable `method` is no longer supported. Pass the name of a " + '`scipy.optimize.minimize` method or "basinhopping".' + ) + from better_optimize.constants import MINIMIZE_MODE_KWARGS + + methods = {k.lower(): k for k in MINIMIZE_MODE_KWARGS} | {"basinhopping": "basinhopping"} + try: + return methods[method.lower()] + except (KeyError, AttributeError): + raise ValueError(f"Unknown method {method!r}. Valid methods are {list(methods.values())}") + + +def _value_vars(vars: Sequence[TensorVariable], model: Model) -> list[Variable]: + try: + value_vars = get_value_vars_from_user_vars(vars, model) + except ValueError as exc: + # Accommodate Deterministics / Potentials by optimizing the free RVs they depend on + value_vars = inputvars(model.replace_rvs_by_values(vars)) + if not value_vars: + raise exc + warnings.warn( + "Intermediate variables (such as Deterministic or Potential) were passed. " + "find_MAP will optimize the underlying free_RVs instead.", + UserWarning, + ) + return value_vars + + +def _unpacked_names(point_map_info, model: Model) -> list[str]: + """One coordinate-aware label per scalar element of the raveled parameter vector.""" + value_to_rv = {value.name: rv.name for value, rv in model.values_to_rvs.items()} + names = [] + for name, shape, *_ in point_map_info: + if not shape: + names.append(name) + continue + dims = model.named_vars_to_dims.get(value_to_rv.get(name, name)) or () + dims = dims if len(dims) == len(shape) else (None,) * len(shape) + axes = [ + coord if (coord := model.coords.get(dim)) is not None and len(coord) == n else range(n) + for dim, n in zip(dims, shape) + ] + names.extend(f"{name}[{','.join(map(str, idx))}]" for idx in product(*axes)) + return names + + +def _fit_dataset(mu: RaveledVars, H_inv: np.ndarray | None, names: list[str]) -> xr.Dataset: + data = {"mean_vector": xr.DataArray(mu.data, dims=["rows"], coords={"rows": names})} + if H_inv is not None: + data["covariance_matrix"] = xr.DataArray( + H_inv, dims=["rows", "columns"], coords={"rows": names, "columns": names} + ) + return xr.Dataset(data) + + +def _optimizer_result_to_dataset( + result: OptimizeResult, method: str, names: list[str] +) -> xr.Dataset: + """Store every field of a scipy ``OptimizeResult``, labelling per-parameter fields by ``names``.""" + from scipy.optimize import LbfgsInvHessProduct, OptimizeResult + + if "lowest_optimization_result" in result: + # basinhopping nests the inner optimizer's result; outer totals (nit, nfev, ...) take precedence + inner = dict(result["lowest_optimization_result"]) + result = OptimizeResult( + inner | {k: v for k, v in result.items() if k != "lowest_optimization_result"} + ) + n = len(names) + data: dict[str, xr.DataArray] = {} + + def add(key, value): + if value is None: + return + if isinstance(value, LbfgsInvHessProduct): + # L-BFGS-B's inverse Hessian is m correction pairs; densifying it is O(n^2) memory + for suffix, pairs in (("sk", value.sk), ("yk", value.yk)): + data[f"{key}_{suffix}"] = xr.DataArray( + np.asarray(pairs), dims=("lbfgs_corrections", "variables") + ) + return + if key == "message": + value = str(value) + try: + value = np.asarray(value) + except ValueError: # ragged tuple, e.g. nelder-mead's (vertices, values) final_simplex + for i, v in enumerate(value): + add(f"{key}_{i}", v) + return + dims = [f"{key}_dim_{i}" for i in range(value.ndim)] + if ( + value.ndim in (1, 2) and value.shape[-1] == n + ): # per-parameter axis, only when sizes match + dims[-1] = "variables" + if value.shape == (n, n): + dims = ["variables", "variables_aux"] + data[key] = xr.DataArray(value, dims=dims) + + for key, value in result.items(): + add(key, value) + data["method"] = xr.DataArray(method) # trust-constr reports its own sub-method under this key + coords = { + d: names + for d in ("variables", "variables_aux") + if any(d in da.dims for da in data.values()) + } + return xr.Dataset(data, coords=coords) + def find_MAP( - start=None, - vars: Sequence[Variable] | None = None, - method="L-BFGS-B", - return_raw=False, - include_transformed=True, - progressbar=True, - progressbar_theme=default_progress_theme, - maxeval=5000, - model=None, - *args, - seed: int | None = None, - **kwargs, -): - """Find the local maximum a posteriori point given a model. + method: minimize_method | Literal["basinhopping"] = "L-BFGS-B", + *, + vars: Sequence[TensorVariable] | None = None, + use_grad: bool | None = None, + use_hess: bool | None = None, + use_hessp: bool | None = None, + initvals: StartDict | None = None, + jitter: bool | None = None, + jitter_max_retries: int = 10, + random_seed: RandomState = None, + progressbar: bool | ProgressBarOptions = True, + compute_hessian: bool = False, + return_inferencedata: bool = False, + idata_kwargs: dict[str, Any] | None = None, + model: Model | None = None, + backend: str | None = None, + compile_kwargs: dict | None = None, + **optimizer_kwargs, +) -> DataTree | PointType: + """Find the local maximum a posteriori point of a model with ``scipy.optimize``. `find_MAP` should not be used to initialize the NUTS sampler. Simply call - ``pymc.sample()`` and it will automatically initialize NUTS in a better - way. + ``pymc.sample()`` and it will automatically initialize NUTS in a better way. Parameters ---------- - start: `dict` of parameter values (Defaults to `model.initial_point`) - These values will be fixed and used for any free RandomVariables that are - not being optimized. - vars: list of TensorVariable - List of free RandomVariables to optimize the posterior with respect to. - Defaults to all continuous RVs in a model. The respective value variables - may also be passed instead. - method: string or callable, optional - Optimization algorithm. Defaults to 'L-BFGS-B' unless discrete variables are - specified in `vars`, then `Powell` which will perform better. For instructions - on use of a callable, refer to SciPy's documentation of `optimize.minimize`. - return_raw: bool, optional defaults to False - Whether to return the full output of scipy.optimize.minimize - include_transformed: bool, optional defaults to True - Flag for reporting automatically unconstrained transformed values in addition - to the constrained values - progressbar: bool, optional defaults to True - Whether to display a progress bar in the command line. - progressbar_theme: Theme, optional - Custom theme for the progress bar. - maxeval: int, optional, defaults to 5000 - The maximum number of times the posterior distribution is evaluated. - model: Model (optional if in `with` context) - *args, **kwargs - Extra args passed to scipy.optimize.minimize - - Notes - ----- - Older code examples used `find_MAP` to initialize the NUTS sampler, - but this is not an effective way of choosing starting values for sampling. - As a result, we have greatly enhanced the initialization of NUTS and - wrapped it inside ``pymc.sample()`` and you should thus avoid this method. + method : str + Optimization method. Any ``scipy.optimize.minimize`` method (Nelder-Mead, Powell, CG, + BFGS, L-BFGS-B, TNC, COBYLA, SLSQP, trust-constr, dogleg, trust-ncg, trust-exact, + trust-krylov, Newton-CG) or ``"basinhopping"``. Defaults to ``"L-BFGS-B"``. + vars : list of TensorVariable, optional + Free random variables (or their value variables) to optimize over. All other variables + are held fixed at their initial values. Defaults to all continuous variables. Passing + discrete variables switches to the gradient-free ``"powell"`` method. + use_grad, use_hess, use_hessp : bool, optional + Whether to compile and pass the gradient, hessian and hessian-vector product to the + optimizer. ``None`` (default) chooses based on ``method``. If gradients are requested + automatically but the model has none, ``"powell"`` is used instead. + initvals : dict, optional + Initial values for (transformed) variables, overriding the model defaults. Partial + initialization is permitted, as in :func:`pymc.sample`. + jitter : bool, optional + Add U(-1, 1) jitter to the initial point of the optimized variables, as ``pymc.sample`` + does. This avoids getting stuck at saddle points of the default initial point (e.g. + products of zero-centered variables). Set ``random_seed`` for reproducible results. + Not jittering is the current default, with a ``FutureWarning``; a future release will + jitter by default. + jitter_max_retries : int + Maximum number of attempts at drawing a jittered initial point with finite log-probability. + random_seed : int, array-like of int, or Generator, optional + Seed for jitter and stochastic optimizers (basinhopping). With a fixed seed the result is + fully reproducible. + progressbar : bool or ProgressBarOptions, default True + Whether to display the optimizer's progress bar. The string options of + :func:`pymc.sample` are accepted and simply enable it. + compute_hessian : bool, default False + Store the inverse Hessian of the negative ``model.logp(jacobian=False)`` at the optimum, + taken over the optimized (unconstrained) value variables, as ``fit.covariance_matrix``. + This needs ``n`` Hessian-vector products and an ``n x n`` matrix, so it is expensive for + large models. + return_inferencedata : bool, default False + If True, return an :class:`arviz.InferenceData` with the MAP point as a single-draw + ``posterior`` (plus ``fit``, ``optimizer_result``, ``observed_data`` and ``constant_data`` + groups). If False, return a ``dict`` mapping variable names to values, transformed ones + included. + + .. deprecated:: + The ``dict`` return is deprecated: a future release will default to True, and + later remove the option. + idata_kwargs : dict, optional + Keyword arguments for :func:`pymc.to_inference_data`, e.g. ``include_transformed=True`` to + also return transformed (unconstrained) values such as ``sigma_log__``. + model : Model (optional if in ``with`` context) + Pass a model from :func:`pymc.model.transform.freeze_model` for constant folding and + compiled functions cached across calls. + backend : str, optional + Computational backend, one of "numba", "c" or "jax". Defaults to the PyTensor default mode. + compile_kwargs : dict, optional + Keyword arguments for the compiled functions. ``compile_kwargs["mode"]`` cannot be combined + with ``backend``. + **optimizer_kwargs + Passed on to ``scipy.optimize.minimize`` (e.g. ``maxiter``, ``tol``), or + ``scipy.optimize.basinhopping`` when ``method="basinhopping"``, in which case + ``minimizer_kwargs["method"]`` selects the inner optimizer (default ``"L-BFGS-B"``). + + Returns + ------- + arviz.InferenceData or dict + MAP estimate, see ``return_inferencedata``. """ - model = modelcontext(model) - - if vars is None: - vars = model.continuous_value_vars - if not vars: - raise ValueError("Model has no unobserved continuous variables.") - else: - try: - vars = get_value_vars_from_user_vars(vars, model) - except ValueError as exc: - # Accommodate case where user passed non-pure RV nodes - vars = inputvars(model.replace_rvs_by_values(vars)) - if vars: - warnings.warn( - "Intermediate variables (such as Deterministic or Potential) were passed. " - "find_MAP will optimize the underlying free_RVs instead.", - UserWarning, - ) - else: - raise exc - - disc_vars = list(typefilter(vars, discrete_types)) - ipfn = make_initial_point_fn( + if isinstance(method, dict): + optimizer_kwargs["start"], method = method, "L-BFGS-B" + explicit = {"initvals": initvals, "random_seed": random_seed} + for old, new in _LEGACY_KWARGS.items(): + if old in optimizer_kwargs: + if new in optimizer_kwargs or explicit.get(new) is not None: + raise ValueError(f"Cannot pass both `{old}` and `{new}`; `{old}` is deprecated.") + warnings.warn( + f"`{old}` is deprecated, use `{new}` instead.", FutureWarning, stacklevel=2 + ) + optimizer_kwargs[new] = optimizer_kwargs.pop(old) + initvals = optimizer_kwargs.pop("initvals", initvals) + random_seed = optimizer_kwargs.pop("random_seed", random_seed) + return_raw = optimizer_kwargs.pop("return_raw", False) + if return_raw: + warnings.warn( + "`return_raw` is deprecated. The optimizer result is stored in the `optimizer_result` " + "group of the returned InferenceData.", + FutureWarning, + stacklevel=2, + ) + if optimizer_kwargs.pop("progressbar_theme", None) is not None: + warnings.warn("`progressbar_theme` is ignored by find_MAP.", FutureWarning, stacklevel=2) + idata_kwargs = {} if idata_kwargs is None else dict(idata_kwargs) + if "include_transformed" in optimizer_kwargs: + if "include_transformed" in idata_kwargs: + raise ValueError("Pass `include_transformed` only via `idata_kwargs`.") + warnings.warn( + "`include_transformed` is deprecated, pass it via `idata_kwargs` as in `pymc.sample`.", + FutureWarning, + stacklevel=2, + ) + idata_kwargs["include_transformed"] = optimizer_kwargs.pop("include_transformed") + + if not return_inferencedata: + warnings.warn( + "`find_MAP` will return an `InferenceData` instead of a dict in a future release. " + "Pass `return_inferencedata=True` to adopt the new behavior now.", + FutureWarning, + stacklevel=2, + ) + if jitter is None: + warnings.warn( + "`find_MAP` will jitter its initial point by default in a future release. Pass " + "`jitter=True` to adopt that now, or `jitter=False` to keep the current behavior.", + FutureWarning, + stacklevel=2, + ) + jitter = False + fit = _fit_MAP( + method, + vars=vars, + use_grad=use_grad, + use_hess=use_hess, + use_hessp=use_hessp, + initvals=initvals, + jitter=jitter, + jitter_max_retries=jitter_max_retries, + random_seed=random_seed, + progressbar=bool(progressbar), + compute_hessian=compute_hessian, model=model, - jitter_rvs=set(), - return_transformed=True, - overrides=start, + backend=backend, + compile_kwargs=compile_kwargs, + **optimizer_kwargs, + ) + out = ( + _map_to_inference_data(fit, idata_kwargs) + if return_inferencedata + else fit.as_point(idata_kwargs.get("include_transformed", True)) ) - start = ipfn(seed) - model.check_start_vals(start) + return (out, fit.res) if return_raw else out # type: ignore[return-value] + + +class _MAPFit(NamedTuple): + model: Model + point: PointType # value variables at the optimum + values: dict[str, np.ndarray] # every unobserved value: free, transformed and deterministic + fn: Function + res: OptimizeResult + x_star: RaveledVars + H_inv: np.ndarray | None + method: str + + def as_point(self, include_transformed: bool) -> PointType: + names = get_default_varnames(self.values, include_transformed) + return {name: self.values[name] for name in names} + + +def _fit_MAP( + method: str, + *, + vars: Sequence[TensorVariable] | None, + use_grad: bool | None, + use_hess: bool | None, + use_hessp: bool | None, + initvals: StartDict | None, + jitter: bool, + jitter_max_retries: int, + random_seed: RandomState, + progressbar: bool, + compute_hessian: bool, + model: Model | None, + backend: str | None, + compile_kwargs: dict | None, + **optimizer_kwargs, +) -> _MAPFit: + """Run the optimization behind :func:`find_MAP`; no defaults, so they live only there.""" + from better_optimize import basinhopping, minimize + + model = cast(Model, modelcontext(model)) + selected = set(model.continuous_value_vars if vars is None else _value_vars(vars, model)) + vars = [var for var in model.value_vars if var in selected] # model order + compile_kwargs = resolve_backend_compile_kwargs(backend, compile_kwargs) + + if not vars: + raise ValueError("Model has no unobserved continuous variables.") + discrete = typefilter(vars, discrete_types) + if compute_hessian and discrete: + raise ValueError( + f"`compute_hessian` is undefined for discrete variables {[v.name for v in discrete]}; " + "exclude them via `vars`." + ) - vars_dict = {var.name: var for var in vars} - x0 = DictToArrayBijection.map( - {var_name: value for var_name, value in start.items() if var_name in vars_dict} + rng = get_random_generator(random_seed) + [ipfn] = make_initial_point_fns_per_chain( + model=model, + overrides=initvals, + jitter_rvs={model.values_to_rvs[v] for v in vars if v not in discrete} if jitter else set(), + chains=1, + ) + start = ipfn(int(rng.integers(2**30))) + + method = _canonical_method(method) + do_basinhopping = method == "basinhopping" + minimizer_kwargs = dict(optimizer_kwargs.pop("minimizer_kwargs", {})) + if do_basinhopping: + method = _canonical_method(minimizer_kwargs.pop("method", "L-BFGS-B")) + + from better_optimize.constants import MINIMIZE_MODE_KWARGS + + auto_grad = use_grad is None + if auto_grad and discrete and MINIMIZE_MODE_KWARGS[method]["uses_grad"]: + warnings.warn( + "Discrete variables are being optimized, so gradients are not available. " + f"Using the gradient-free method 'powell' instead of '{method}'.", + UserWarning, + ) + method, use_grad = "powell", False + use_grad, use_hess, use_hessp = set_optimizer_function_defaults( + method, use_grad, use_hess, use_hessp ) - # TODO: If the mapping is fixed, we can simply create graphs for the - # mapping and avoid all this bijection overhead - compiled_logp_func = DictToArrayBijection.mapf(model.compile_logp(jacobian=False), start) - logp_func = lambda x: compiled_logp_func(RaveledVars(x, x0.point_map_info)) # noqa: E731 + # Held-fixed variables are shared inputs, so one compiled loss serves every candidate start + fixed = { + v: pytensor.shared(start[v.name], v.name, shape=v.type.shape) + for v in model.value_vars + if v not in vars + } + loss = -cast(TensorVariable, model.logp(jacobian=False)) + if fixed: + loss = graph_replace(loss, fixed) + + def compile_funcs(use_grad, use_hess, use_hessp): + return scipy_optimize_funcs_from_loss( + loss=loss, + inputs=vars, + initial_point_dict=start, + use_grad=use_grad, + use_hess=use_hess, + use_hessp=use_hessp, + compile_kwargs=compile_kwargs, + ) - rvs = [model.values_to_rvs[vars_dict[name]] for name, _, _, _ in x0.point_map_info] + # Compile the hessp compute_hessian needs alongside the loss, rather than recompiling the loss later + need_hessp = compute_hessian and not use_hess try: - # This might be needed for calls to `dlogp_func` - # start_map_info = tuple((v.name, v.shape, v.dtype) for v in vars) - compiled_dlogp_func = DictToArrayBijection.mapf( - model.compile_dlogp(rvs, jacobian=False), start + f_fused, f_hessp = compile_funcs(use_grad, use_hess, use_hessp or need_hessp) + except (NotImplementedError, tg.NullTypeGradError) as exc: + if not (auto_grad and use_grad): + raise + warnings.warn( + f"Gradient not available ({exc}). Using the gradient-free method 'powell' instead of " + f"'{method}'.", + UserWarning, ) - dlogp_func = lambda x: compiled_dlogp_func(RaveledVars(x, x0.point_map_info)) # noqa: E731 - compute_gradient = True - except (AttributeError, NotImplementedError, tg.NullTypeGradError): - compute_gradient = False - - if disc_vars or not compute_gradient: - pm._log.warning( - "Warning: gradient not available." - + "(E.g. vars contains discrete variables). MAP " - + "estimates may not be accurate for the default " - + "parameters. Defaulting to non-gradient minimization " - + "'Powell'." + method, use_grad, use_hess, use_hessp = "powell", False, False, False + f_fused, f_hessp = compile_funcs(use_grad, use_hess, need_hessp) + + def ravel(point): + return DictToArrayBijection.map({str(v.name): point[str(v.name)] for v in vars}) + + def loss_at(point): + for var, shared in fixed.items(): + shared.set_value(point[var.name]) + out = f_fused(ravel(point).data) + return out[0] if isinstance(out, tuple | list) else out + + # As in pm.sample, redraw a jittered start until the objective is finite there + for _ in range(jitter_max_retries if jitter else 0): + if np.isfinite(loss_at(start)): + break + start = ipfn(int(rng.integers(2**30))) + else: + if not np.isfinite(loss_at(start)): + model.check_start_vals(start) + x0 = ravel(start) + + optimizer_hessp = f_hessp if use_hessp else None + if do_basinhopping: + minimizer_kwargs = {"method": method, "hessp": optimizer_hessp, **minimizer_kwargs} + optimizer_kwargs.setdefault("rng", rng) + res = basinhopping( + func=f_fused, + x0=x0.data, + progressbar=progressbar, + minimizer_kwargs=minimizer_kwargs, + **optimizer_kwargs, ) - method = "Powell" - - if compute_gradient and method != "Powell": - cost_func = CostFuncWrapper(maxeval, progressbar, progressbar_theme, logp_func, dlogp_func) else: - cost_func = CostFuncWrapper(maxeval, progressbar, progressbar_theme, logp_func) - compute_gradient = False + res = minimize( + f=f_fused, + x0=x0.data, + hessp=optimizer_hessp, + progressbar=progressbar, + method=method, + **optimizer_kwargs, + ) - with cost_func.progress: - try: - opt_result = optimize.minimize( - cost_func, x0.data, method=method, jac=compute_gradient, *args, **kwargs - ) - mx0 = opt_result["x"] # r -> opt_result - except (KeyboardInterrupt, StopIteration) as e: - mx0, opt_result = cost_func.previous_x, None - if isinstance(e, StopIteration): - pm._log.info(e) - finally: - cost_func.progress.update(cost_func.task, completed=cost_func.n_eval, refresh=True) - - mx0 = RaveledVars(mx0, x0.point_map_info) - unobserved_vars = get_default_varnames(model.unobserved_value_vars, include_transformed) - unobserved_vars_values = model.compile_fn(inputs=model.value_vars, outs=unobserved_vars)( - DictToArrayBijection.rmap(mx0, start) + if not res.get("success", True): + warnings.warn(f"The optimizer did not converge: {res.get('message', '')}", UserWarning) + + H_inv = _compute_inverse_hessian(res.x, f_fused, f_hessp, use_hess) if compute_hessian else None + x_star = RaveledVars(np.asarray(res.x), x0.point_map_info) + point = DictToArrayBijection.rmap(x_star, start) + + # Free, transformed and deterministic values at the optimum; one-shot, so avoid a heavy compile + unobserved = model.unobserved_value_vars + fn = model.compile_fn( + unobserved, + inputs=model.value_vars, + point_fn=False, + on_unused_input="ignore", + mode="FAST_COMPILE", ) - mx = {var.name: value for var, value in zip(unobserved_vars, unobserved_vars_values)} + values = dict(zip([var.name for var in unobserved], fn(**point))) - if return_raw: - return mx, opt_result - else: - return mx - - -def allfinite(x): - return np.all(isfinite(x)) - - -class CostFuncWrapper: - def __init__( - self, - maxeval=5000, - progressbar=True, - progressbar_theme=default_progress_theme, - logp_func=None, - dlogp_func=None, - ): - self.n_eval = 0 - self.maxeval = maxeval - self.logp_func = logp_func - if dlogp_func is None: - self.use_gradient = False - self.desc = "logp = {:,.5g}" - else: - self.dlogp_func = dlogp_func - self.use_gradient = True - self.desc = "logp = {:,.5g}, ||grad|| = {:,.5g}" - self.previous_x = None - self.progressbar = progressbar - self.progress = CustomProgress( - *Progress.get_default_columns(), - TextColumn("{task.fields[loss]}"), - console=Console(theme=progressbar_theme), - disable=not progressbar, - ) - self.task = self.progress.add_task("MAP", total=maxeval, loss="") - - def __call__(self, x): - neg_value = np.float64(self.logp_func(floatX(x))) - value = -1.0 * neg_value - if self.use_gradient: - neg_grad = self.dlogp_func(floatX(x)) - if np.all(np.isfinite(neg_grad)): - self.previous_x = x - grad = -1.0 * neg_grad - grad = grad.astype(np.float64) - else: - self.previous_x = x - grad = None - - if self.n_eval % 10 == 0: - self.progress.update(self.task, loss=self.update_progress_desc(neg_value, grad)) - - if self.n_eval > self.maxeval: - self.progress.update(self.task, loss=self.update_progress_desc(neg_value, grad)) - raise StopIteration - - self.n_eval += 1 - self.progress.update(self.task, completed=self.n_eval) - - if self.use_gradient: - return value, grad - else: - return value - - def update_progress_desc(self, neg_value: float, grad: np.float64 = None) -> None: - if self.progressbar: - if grad is None: - return self.desc.format(neg_value) - else: - norm_grad = np.linalg.norm(grad) - return self.desc.format(neg_value, norm_grad) + return _MAPFit( + model, point, values, fn, res, x_star, H_inv, "basinhopping" if do_basinhopping else method + ) + + +def _find_MAP_point( + *, + model: Model, + initvals: StartDict | None, + jitter: bool, + jitter_max_retries: int, + random_seed: RandomState, + progressbar: bool, + compile_kwargs: dict | None, +) -> PointType: + """MAP point, transformed values included, for internal callers like ``init="jitter+map"``.""" + fit = _fit_MAP( + "L-BFGS-B", + vars=None, + use_grad=None, + use_hess=None, + use_hessp=None, + initvals=initvals, + jitter=jitter, + jitter_max_retries=jitter_max_retries, + random_seed=random_seed, + progressbar=progressbar, + compute_hessian=False, + model=model, + backend=None, + compile_kwargs=compile_kwargs, + ) + return fit.as_point(include_transformed=True) + + +def _map_to_inference_data(fit: _MAPFit, idata_kwargs: dict[str, Any]) -> DataTree: + model, values = fit.model, fit.values + trace = NDArray( + model=model, + fn=fit.fn, + var_shapes={k: v.shape for k, v in values.items()}, + var_dtypes={k: v.dtype for k, v in values.items()}, + ) + trace.setup(draws=1, chain=0) + trace.record(fit.point, in_warmup=False) + trace.close() + idata = to_inference_data(MultiTrace([trace]), model=model, **idata_kwargs) + if not idata["sample_stats"].data_vars: # a single optimum has no sampler stats + del idata["sample_stats"] + labels = _unpacked_names(fit.x_star.point_map_info, model) + idata["fit"] = DataTree(dataset=_fit_dataset(fit.x_star, fit.H_inv, labels)) + idata["optimizer_result"] = DataTree( + dataset=_optimizer_result_to_dataset(fit.res, fit.method, labels) + ) + return idata diff --git a/pyproject.toml b/pyproject.toml index 882f4e9f97..4a97214bf0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -33,6 +33,7 @@ keywords = [ ] dependencies = [ "arviz>=1.1.0,<2.0", + "better-optimize>=0.4.2,<0.5", "cachetools>=4.2.1,<7", "cloudpickle", "numpy>=1.25.0", diff --git a/requirements-dev.txt b/requirements-dev.txt index 442037037b..d97b480609 100644 --- a/requirements-dev.txt +++ b/requirements-dev.txt @@ -2,6 +2,7 @@ # See that file for comments about the need/usage of each dependency. arviz>=1.1.0,<2.0 +better-optimize>=0.4.2,<0.5 cachetools>=4.2.1,<7 cloudpickle ipython>=7.16 diff --git a/tests/distributions/test_truncated.py b/tests/distributions/test_truncated.py index 72b21264b1..44de2287df 100644 --- a/tests/distributions/test_truncated.py +++ b/tests/distributions/test_truncated.py @@ -434,9 +434,13 @@ def test_truncated_inference(): observed=obs, ) - map = find_MAP(progressbar=False) + lam_map = ( + find_MAP(progressbar=False, jitter=False, return_inferencedata=True) + .posterior["lam"] + .item() + ) - assert np.isclose(map["lam"], lam_true, atol=0.1) + assert np.isclose(lam_map, lam_true, atol=0.1) def test_truncated_gamma(): diff --git a/tests/gp/test_gp.py b/tests/gp/test_gp.py index e992c38009..badd0ec5b8 100644 --- a/tests/gp/test_gp.py +++ b/tests/gp/test_gp.py @@ -41,7 +41,11 @@ def setup_method(self): self.gp = pm.gp.Marginal(mean_func=mean_func, cov_func=cov_func) sigma = pm.HalfNormal("sigma", sigma=100) self.gp.marginal_likelihood("lik", self.x[:, None], self.y, sigma) - self.map_full = pm.find_MAP(method="bfgs") # bfgs seems to work much better than lbfgsb + # bfgs seems to work much better than lbfgsb + posterior = pm.find_MAP( + method="bfgs", jitter=False, return_inferencedata=True + ).posterior.isel(chain=0, draw=0) + self.map_full = {name: da.values for name, da in posterior.items()} self.x_new = np.linspace(-6, 6, 20) @@ -71,7 +75,10 @@ def test_fits_and_preds(self, approx): gp = pm.gp.MarginalApprox(mean_func=mean_func, cov_func=cov_func, approx=approx) sigma = pm.HalfNormal("sigma", sigma=100, initval=50.0) gp.marginal_likelihood("lik", self.x[:, None], self.x[:, None], self.y, sigma) - map_approx = pm.find_MAP(method="bfgs") + posterior = pm.find_MAP( + method="bfgs", jitter=False, return_inferencedata=True + ).posterior.isel(chain=0, draw=0) + map_approx = {name: da.values for name, da in posterior.items()} # Check MAP gets approximately correct result npt.assert_allclose(self.map_full["c"], map_approx["c"], atol=0.01, rtol=0.1) diff --git a/tests/model/transform/test_conditioning.py b/tests/model/transform/test_conditioning.py index d4c2235135..a083170837 100644 --- a/tests/model/transform/test_conditioning.py +++ b/tests/model/transform/test_conditioning.py @@ -337,15 +337,19 @@ def test_change_value_transforms(): new_p = transformed_p["p"] assert transformed_p.rvs_to_transforms[new_p] == logodds assert transformed_p.rvs_to_values[new_p].name == "p_logodds__" - mean_q = pm.find_MAP(progressbar=False) + mean_q = ( + pm.find_MAP(progressbar=False, jitter=False, return_inferencedata=True) + .posterior["p"] + .item() + ) with change_value_transforms(transformed_p, {"p": None}) as untransformed_p: new_p = untransformed_p["p"] assert untransformed_p.rvs_to_transforms[new_p] is None assert untransformed_p.rvs_to_values[new_p].name == "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] - np.testing.assert_allclose(np.round(mean_q["p"], 2), 0.67) + np.testing.assert_allclose(np.round(mean_q, 2), 0.67) np.testing.assert_allclose(np.round(std_q[0], 2), 0.16) diff --git a/tests/sampling/test_mcmc.py b/tests/sampling/test_mcmc.py index dd17a193ba..f7b383e3dd 100644 --- a/tests/sampling/test_mcmc.py +++ b/tests/sampling/test_mcmc.py @@ -64,7 +64,7 @@ class TestSample: def setup_method(self): self.model, self.start, self.step, _ = simple_init() - @pytest.mark.parametrize("init", ("jitter+adapt_diag", "advi", "map")) + @pytest.mark.parametrize("init", ("jitter+adapt_diag", "advi", "jitter+map")) @pytest.mark.parametrize("cores", (1, 2)) @pytest.mark.parametrize( "chains, seeds", @@ -160,6 +160,7 @@ def test_sample_does_not_rely_on_external_global_seeding(self): "advi", "advi_map", "map", + "jitter+map", "adapt_diag", "jitter+adapt_diag", "jitter+adapt_diag_grad", @@ -586,7 +587,7 @@ def test_sample_find_MAP_does_not_modify_start(): # make sure find_Map does not modify the start dict start = {"untransformed": 2} - pm.find_MAP(start=start) + pm.find_MAP(initvals=start, jitter=False, return_inferencedata=True, progressbar=False) assert start == {"untransformed": 2} # make sure sample does not modify the start dict @@ -703,6 +704,7 @@ def check_exec_nuts_init(method): "jitter+adapt_diag", "adapt_diag", "map", + "jitter+map", "adapt_full", "jitter+adapt_full", ], @@ -715,6 +717,26 @@ def test_exec_nuts_init(method): check_exec_nuts_init(method) +def test_init_map_jitter_and_initvals(monkeypatch): + calls = [] + find_map_point = pm.sampling.mcmc._find_MAP_point + monkeypatch.setattr( + pm.sampling.mcmc, + "_find_MAP_point", + lambda **kwargs: calls.append(kwargs) or find_map_point(**kwargs), + ) + for init in ("map", "jitter+map", "advi_map"): + check_exec_nuts_init(init) + assert [c["jitter"] for c in calls] == [False, False, True, True, False, False] + + calls.clear() + with pm.Model(): + pm.Normal("a") + pm.init_nuts(init="map", initvals={"a": 2.0}, random_seed=[1]) + pm.init_nuts(init="map", initvals=[{"a": 3.0}, {"a": 4.0}], chains=2, random_seed=[1, 2]) + assert [c["initvals"] for c in calls] == [{"a": 2.0}, {"a": 3.0}] + + @pytest.mark.skip(reason="Test requires monkey patching of RandomGenerator") @pytest.mark.parametrize( "initval, jitter_max_retries, expectation", diff --git a/tests/tuning/test_scipy_interface.py b/tests/tuning/test_scipy_interface.py new file mode 100644 index 0000000000..8e985b2ddd --- /dev/null +++ b/tests/tuning/test_scipy_interface.py @@ -0,0 +1,171 @@ +# Copyright 2024 - present The PyMC Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import numpy as np +import pytest + +from pytensor import tensor as pt + +from pymc.tuning import scipy_interface +from pymc.tuning.scipy_interface import ( + scipy_optimize_funcs_from_loss, + set_optimizer_function_defaults, +) + + +@pytest.fixture +def simple_loss_and_inputs(): + x = pt.vector("x") + return pt.sum(x**2), [x] + + +def test_set_optimizer_function_defaults_warns_and_prefers_hessp(): + with pytest.warns(UserWarning, match="Only one of `use_hess` and `use_hessp`"): + flags = set_optimizer_function_defaults("trust-ncg", True, True, True) + assert flags == (True, False, True) + + +@pytest.mark.parametrize( + "method, use_hess, use_hessp, expected", + [ + ("trust-ncg", None, None, (True, False, True)), + ("trust-ncg", None, True, (True, False, True)), + ("trust-ncg", True, None, (True, True, False)), + ("trust-ncg", False, None, (True, False, True)), + ("L-BFGS-B", None, None, (True, False, False)), + # setting one flag must not flip the other on for a method that cannot use it + ("L-BFGS-B", None, False, (True, False, False)), + ("L-BFGS-B", False, None, (True, False, False)), + ("trust-exact", None, False, (True, True, False)), + ("powell", None, None, (False, False, False)), + ], +) +def test_set_optimizer_function_defaults(method, use_hess, use_hessp, expected): + assert set_optimizer_function_defaults(method, None, use_hess, use_hessp) == expected + + +@pytest.mark.parametrize( + "method, flags, expected", + [ + ("trust-exact", (None, None, True), (True, True, False)), # keeps the Hessian it needs + ("L-BFGS-B", (None, True, None), (True, False, False)), + ("powell", (True, None, None), (False, False, False)), + ], +) +def test_set_optimizer_function_defaults_ignores_unused_flags(method, flags, expected): + with pytest.warns(UserWarning, match=f"Method '{method}' does not use"): + assert set_optimizer_function_defaults(method, *flags) == expected + + +@pytest.mark.parametrize( + "use_grad, use_hess, use_hessp", + [(False, False, False), (True, False, False), (True, True, False), (True, False, True)], +) +def test_scipy_optimize_funcs_from_loss(simple_loss_and_inputs, use_grad, use_hess, use_hessp): + loss, inputs = simple_loss_and_inputs + f_fused, f_hessp = scipy_optimize_funcs_from_loss( + loss, + inputs, + use_grad=use_grad, + use_hess=use_hess, + use_hessp=use_hessp, + inputs_are_flat=True, + ) + x = np.array([1.0, 2.0]) + if not use_grad: + assert np.isclose(f_fused(x), 5.0) + return + loss_val, grad_val, *rest = f_fused(x) + assert np.isclose(loss_val, 5.0) + np.testing.assert_allclose(grad_val, 2 * x) + if use_hess: + np.testing.assert_allclose(rest[0], 2 * np.eye(2)) + if use_hessp: + np.testing.assert_allclose(f_hessp(x, np.array([1.0, 0.0])), [2.0, 0.0]) + else: + assert f_hessp is None + + +def test_scipy_optimize_funcs_from_loss_hess_without_grad(simple_loss_and_inputs): + loss, inputs = simple_loss_and_inputs + with pytest.raises(ValueError, match="Cannot compute hessian without"): + scipy_optimize_funcs_from_loss( + loss, inputs, {"x": np.zeros(2)}, use_grad=False, use_hess=True + ) + + +def test_scipy_optimize_funcs_from_loss_flat_input(simple_loss_and_inputs): + loss, [x] = simple_loss_and_inputs + f_fused, _ = scipy_optimize_funcs_from_loss(loss, x, use_grad=True, inputs_are_flat=True) + loss_val, grad_val = f_fused(np.array([1.0, 2.0])) + assert np.isclose(loss_val, 5.0) + np.testing.assert_allclose(grad_val, [2.0, 4.0]) + + +@pytest.mark.parametrize("use_hess", [True, False]) +def test_compute_inverse_hessian_is_exact(use_hess): + x = pt.vector("x", shape=(2,)) + A = np.array([[3.0, 1.0], [1.0, 2.0]]) + f_fused, f_hessp = scipy_optimize_funcs_from_loss( + loss=0.5 * x @ A @ x, + inputs=[x], + initial_point_dict={"x": np.zeros(2)}, + use_grad=True, + use_hess=use_hess, + use_hessp=not use_hess, + ) + H_inv = scipy_interface._compute_inverse_hessian(np.ones(2), f_fused, f_hessp, use_hess) + np.testing.assert_allclose(H_inv, np.linalg.inv(A)) + + +def test_compute_inverse_hessian_indefinite(): + x = pt.vector("x", shape=(2,)) + A = np.array([[1.0, 0.0], [0.0, -1.0]]) # saddle point: not a minimum + _, f_hessp = scipy_optimize_funcs_from_loss( + loss=0.5 * x @ A @ x, + inputs=[x], + initial_point_dict={"x": np.zeros(2)}, + use_grad=True, + use_hessp=True, + ) + with pytest.warns(UserWarning, match="not positive definite"): + H_inv = scipy_interface._compute_inverse_hessian(np.zeros(2), f_hessp=f_hessp) + assert np.all(np.isfinite(H_inv)) + np.testing.assert_allclose(H_inv[0, 0], 1.0) # the well-defined direction is untouched + assert np.all(np.linalg.eigvalsh(H_inv) > 0) + + +def test_compute_inverse_hessian_requires_second_order(): + with pytest.raises(ValueError, match="Either `f_hessp`"): + scipy_interface._compute_inverse_hessian(np.zeros(2)) + + +def test_scipy_optimize_funcs_from_loss_jax_mode(): + pytest.importorskip("jax") + x = pt.tensor("x", shape=(2,)) + loss = (x[0] ** 2 + 2) + (x[0] * x[1] + 3) + f_fused, f_hessp = scipy_optimize_funcs_from_loss( + loss=loss, + inputs=[x], + initial_point_dict={"x": np.array([1.0, 2.0])}, + use_grad=True, + use_hess=True, + use_hessp=True, + compile_kwargs={"mode": "JAX"}, + ) + x_val = np.array([1.0, 2.0]) + z, grad, hess = f_fused(x_val) + np.testing.assert_allclose(z, 8.0) + np.testing.assert_allclose(np.asarray(grad).squeeze(), [2 * x_val[0] + x_val[1], x_val[0]]) + np.testing.assert_allclose(np.asarray(hess).squeeze(), [[2, 1], [1, 0]]) + np.testing.assert_allclose(np.asarray(f_hessp(x_val, np.array([1.0, 0.0]))).squeeze(), [2, 1]) diff --git a/tests/tuning/test_starting.py b/tests/tuning/test_starting.py index 0b19aa57ff..7973c1a259 100644 --- a/tests/tuning/test_starting.py +++ b/tests/tuning/test_starting.py @@ -12,21 +12,55 @@ # See the License for the specific language governing permissions and # limitations under the License. import re +import warnings + +from functools import partial import numpy as np +import pytensor +import pytensor.tensor as pt import pytest +import xarray as xr from numpy.testing import assert_allclose +from scipy.optimize import LbfgsInvHessProduct, OptimizeResult import pymc as pm -from pymc.exceptions import ImputationWarning +from pymc.exceptions import ImputationWarning, SamplingError +from pymc.model.transform.optimization import freeze_model from pymc.step_methods.metropolis import tune -from pymc.testing import select_by_precision -from pymc.tuning import find_MAP +from pymc.testing import fast_unstable_sampling_mode, select_by_precision +from pymc.tuning.starting import _find_MAP_point, _optimizer_result_to_dataset from tests import models from tests.models import non_normal, simple_arbitrary_det, simple_model +# The future defaults; the current, deprecated ones are tested through `pm.find_MAP` +find_MAP = partial(pm.find_MAP, return_inferencedata=True, jitter=True) + + +@pytest.fixture(autouse=True) +def fast_compile_mode(): + """Cheap compilation; the default backend is covered by the init="map", GP and JAX tests.""" + with pytensor.config.change_flags(mode=fast_unstable_sampling_mode): + yield + + +def map_point(*args, **kwargs): + """MAP point read from the returned InferenceData's single-draw posterior.""" + posterior = find_MAP(*args, **{"progressbar": False, **kwargs}).posterior + return {name: da.values[0, 0] for name, da in posterior.items()} + + +@pytest.fixture +def normal_model(): + rng = np.random.default_rng(sum(map(ord, "find_MAP"))) + with pm.Model() as m: + mu = pm.Normal("mu") + sigma = pm.Exponential("sigma", 1) + pm.Normal("y_hat", mu=mu, sigma=sigma, observed=rng.normal(loc=3, scale=1.5, size=10)) + return m + @pytest.mark.parametrize("bounded", [False, True]) def test_mle_jacobian(bounded): @@ -34,9 +68,8 @@ def test_mle_jacobian(bounded): truth = 10.0 # Simple normal model should give mu=10.0 rtol = 1e-4 # this rtol should work on both floatX precisions - start, model, _ = models.simple_normal(bounded_prior=bounded) - with model: - map_estimate = find_MAP(method="BFGS", model=model) + _, model, _ = models.simple_normal(bounded_prior=bounded) + map_estimate = map_point(method="BFGS", model=model) assert_allclose(map_estimate["mu_i"], truth, rtol=rtol) @@ -50,7 +83,7 @@ def test_tune_not_inplace(): def test_accuracy_normal(): _, model, (mu, _) = simple_model() with model: - newstart = find_MAP(pm.Point(x=[-10.5, 100.5])) + newstart = map_point(initvals=pm.Point(x=[-10.5, 100.5])) assert_allclose( newstart["x"], [mu, mu], atol=select_by_precision(float64=1e-5, float32=1e-4) ) @@ -59,7 +92,7 @@ def test_accuracy_normal(): def test_accuracy_non_normal(): _, model, (mu, _) = non_normal(4) with model: - newstart = find_MAP(pm.Point(x=[0.5, 0.01, 0.95, 0.99])) + newstart = map_point(initvals=pm.Point(x=[0.5, 0.01, 0.95, 0.99]), jitter=False) assert_allclose(newstart["x"], mu, atol=select_by_precision(float64=1e-5, float32=1e-4)) @@ -76,19 +109,33 @@ def test_find_MAP_discrete(): pm.Binomial("ss", n=n, p=p) pm.Binomial("s", n=n, p=p, observed=yes) - map_est1 = find_MAP() - map_est2 = find_MAP(vars=model.value_vars) + map_est1 = map_point(random_seed=1) + # Joint discrete + continuous powell search: reference values are from the unjittered start + with pytest.warns(UserWarning, match="Discrete variables are being optimized"): + map_est2 = map_point(vars=model.value_vars, jitter=False) - assert_allclose(map_est1["p"], 0.6086956533498806, atol=tol1, rtol=0) + # ss is held fixed at its (jittered-p dependent) initial value; conjugate MAP given ss + ss0 = map_est1["ss"] + assert_allclose( + map_est1["p"], (alpha + yes + ss0 - 1) / (alpha + beta + 2 * n - 2), atol=tol1, rtol=0 + ) assert_allclose(map_est2["p"], 0.695642178810167, atol=tol2, rtol=0) assert map_est2["ss"] == 14 + # Gradient-free methods handle discrete variables themselves, so no fallback to powell + with model, warnings.catch_warnings(): + warnings.simplefilter("error", UserWarning) + idata = find_MAP("nelder-mead", vars=model.value_vars, progressbar=False, random_seed=1) + assert idata.optimizer_result["method"].item() == "nelder-mead" + def test_find_MAP_no_gradient(): _, model = simple_arbitrary_det() - with model: - find_MAP() + with pytest.warns(UserWarning, match="Gradient not available"): + find_MAP(model=model, progressbar=False) + with pytest.raises(NotImplementedError): + find_MAP(model=model, use_grad=True, progressbar=False) def test_find_MAP(): @@ -104,9 +151,9 @@ def test_find_MAP(): pm.Normal("y", mu=mu, tau=sigma**-2, observed=data) # Test gradient minimization - map_est1 = find_MAP(progressbar=False) - # Test non-gradient minimization - map_est2 = find_MAP(progressbar=False, method="Powell") + map_est1 = map_point() + # Test non-gradient minimization, with case-insensitive method name + map_est2 = map_point(method="Powell") assert_allclose(map_est1["mu"], 0, atol=tol) assert_allclose(map_est1["sigma"], 1, atol=tol) @@ -131,8 +178,8 @@ def test_find_MAP_issue_5923(): pm.Normal("y", mu=mu, tau=sigma**-2, observed=data) start = {"mu": -0.5, "sigma": 1.25} - map_est1 = find_MAP(progressbar=False, vars=[mu, sigma], start=start) - map_est2 = find_MAP(progressbar=False, vars=[sigma, mu], start=start) + map_est1 = map_point(vars=[mu, sigma], initvals=start) + map_est2 = map_point(vars=[sigma, mu], initvals=start) assert_allclose(map_est1["mu"], 0, atol=tol) assert_allclose(map_est1["sigma"], 1, atol=tol) @@ -147,7 +194,7 @@ def test_find_MAP_issue_4488(): with pytest.warns(ImputationWarning): x = pm.Gamma("x", alpha=3, beta=10, observed=np.array([1, np.nan])) y = pm.Deterministic("y", x + 1) - map_estimate = find_MAP() + map_estimate = map_point(idata_kwargs={"include_transformed": True}) assert not set.difference({"x_unobserved", "x_unobserved_log__", "y"}, set(map_estimate.keys())) assert_allclose(map_estimate["x_unobserved"], 0.2, rtol=1e-4, atol=1e-4) @@ -163,5 +210,347 @@ def test_find_MAP_warning_non_free_RVs(): msg = "Intermediate variables (such as Deterministic or Potential) were passed" with pytest.warns(UserWarning, match=re.escape(msg)): - r = pm.find_MAP(vars=[det]) + r = map_point(vars=[det], jitter=False) assert_allclose([r["x"], r["y"], r["det"]], [50, 50, 100]) + + +def test_find_MAP_frozen_model(normal_model): + kwargs = {"progressbar": False, "random_seed": 1} + frozen = find_MAP(model=freeze_model(normal_model), **kwargs).posterior["mu"] + # constant folding differs, so the optimizers stop at slightly different points + assert_allclose(frozen, find_MAP(model=normal_model, **kwargs).posterior["mu"], rtol=1e-4) + + +def test_find_MAP_vars_subset_holds_others_fixed(normal_model): + with normal_model: + r = map_point(vars=[normal_model["mu"]], initvals={"sigma": 2.0}) + assert_allclose(r["sigma"], 2.0) + + +@pytest.mark.parametrize( + "method, use_grad, use_hess, use_hessp, compute_hessian", + [ + # compute_hessian once per Hessian route: fused, hessp, and compiled after optimizing + ("Newton-CG", True, True, False, True), + ("Newton-CG", True, False, True, True), + ("L-BFGS-B", True, False, False, True), + ("powell", False, False, False, True), + ("BFGS", True, False, False, False), + ("trust-exact", True, True, False, False), + ("trust-constr", True, False, True, False), + ("nelder-mead", False, False, False, False), + ], +) +def test_find_MAP_inferencedata( + normal_model, method, use_grad, use_hess, use_hessp, compute_hessian +): + include_transformed = compute_hessian + idata = find_MAP( + method=method, + model=normal_model, + use_grad=use_grad, + use_hess=use_hess, + use_hessp=use_hessp, + progressbar=False, + idata_kwargs={"include_transformed": include_transformed}, + compute_hessian=compute_hessian, + ) + assert set(idata.children) == {"posterior", "fit", "optimizer_result", "observed_data"} + + 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 + assert idata.fit.rows.values.tolist() == ["mu", "sigma_log__"] + assert idata.optimizer_result["method"].item() == method + assert ("hess_inv" in idata.optimizer_result) == (method == "BFGS") + assert ("hess_inv_sk" in idata.optimizer_result) == (method == "L-BFGS-B") + for key in ("hess", "hess_inv"): + if key in idata.optimizer_result: + assert idata.optimizer_result[key].dims == ("variables", "variables_aux") + if compute_hessian: + # Exact inverse Hessian of the loss at the optimum, never an optimizer's approximation + mean = idata.fit.mean_vector.values + d2loss = normal_model.compile_d2logp(jacobian=False, negate_output=True) + H = d2loss({"mu": mean[0], "sigma_log__": mean[1]}) + assert_allclose(idata.fit.covariance_matrix.values, np.linalg.inv(H), rtol=1e-6) + + +def test_find_MAP_compute_hessian_discrete_raises(): + with pm.Model() as m: + p = pm.Beta("p", 2, 2) + pm.Binomial("k", n=10, p=p) + with pytest.raises(ValueError, match=r"undefined for discrete variables \['k'\]"): + find_MAP(model=m, vars=m.value_vars, compute_hessian=True, progressbar=False) + + +def test_find_MAP_compute_hessian_float32(): + with pytensor.config.change_flags(floatX="float32"): + with pm.Model() as m: + mu = pm.Normal("mu") + sigma = pm.Exponential("sigma", 1) + pm.Normal("y", mu, sigma, observed=np.linspace(1, 5, 10)) + idata = find_MAP(model=m, compute_hessian=True, progressbar=False, random_seed=1) + assert np.all(np.isfinite(idata.fit.covariance_matrix.values)) + + +@pytest.mark.parametrize("method", ["L-BFGS-B", "powell"]) +def test_find_MAP_jax_backend(normal_model, method): + pytest.importorskip("jax") + idata = find_MAP( + method, model=normal_model, backend="jax", compute_hessian=True, progressbar=False + ) + assert idata.fit.covariance_matrix.shape == (2, 2) + assert_allclose(idata.posterior["mu"].item(), 3.0, atol=1.0) + + +def test_find_MAP_return_inferencedata_consistent(normal_model): + kwargs = {"model": normal_model, "progressbar": False, "random_seed": 1} + idata_kwargs = {"include_transformed": True} + idata = find_MAP(idata_kwargs=idata_kwargs, **kwargs) + with pytest.warns(FutureWarning, match="will return an `InferenceData`"): + point = find_MAP(return_inferencedata=False, idata_kwargs=idata_kwargs, **kwargs) + assert set(point) == {"mu", "sigma", "sigma_log__"} + internal = _find_MAP_point( + initvals=None, jitter=True, jitter_max_retries=10, compile_kwargs=None, **kwargs + ) + assert_allclose(internal["sigma_log__"], point["sigma_log__"]) + for name, value in point.items(): + assert_allclose(idata.posterior[name].values.squeeze(), value) + + +def test_find_MAP_idata_kwargs(normal_model): + idata = find_MAP(model=normal_model, idata_kwargs={"log_likelihood": True}, progressbar=False) + assert "log_likelihood" in idata.children + assert idata.log_likelihood["y_hat"].shape == (1, 1, 10) + + +def test_find_MAP_shared_variables(): + x_val = np.linspace(-1, 1, 20) + with pm.Model() as m: + x = pm.Data("x", x_val) + beta = pm.Normal("beta") + sigma = pm.Exponential("sigma", 1) + pm.Normal( + "y", beta * x, sigma, observed=2 * x_val + np.random.default_rng(0).normal(0, 0.1, 20) + ) + + idata = find_MAP(model=m, progressbar=False) + assert "x" in idata.constant_data + assert "y" in idata.observed_data + assert_allclose(idata.posterior["beta"].item(), 2.0, atol=0.1) + + +@pytest.mark.parametrize("use_hess, use_hessp", [(True, False), (False, True)]) +def test_find_MAP_basinhopping(normal_model, use_hess, use_hessp): + idata = find_MAP( + method="basinhopping", + model=normal_model, + use_hess=use_hess, + use_hessp=use_hessp, + progressbar=False, + random_seed=1, + minimizer_kwargs={"method": "Newton-CG"}, + niter=3, + ) + assert idata.posterior["mu"].shape == (1, 1) + assert idata.optimizer_result["method"].item() == "basinhopping" + assert idata.optimizer_result["nit"].item() == 3 # basinhopping's totals, not one inner run's + + +def test_find_MAP_with_coords(): + with pm.Model(coords={"group": [1, 2, 3, 4, 5]}) as m: + mu_loc = pm.Normal("mu_loc", 0, 1) + mu_scale = pm.HalfNormal("mu_scale", 1) + mu = pm.Normal("mu", mu_loc, mu_scale, dims=["group"]) + sigma = pm.HalfNormal("sigma", 1, dims=["group"]) + pm.Normal("obs", mu=mu, sigma=sigma, observed=np.random.normal(size=(10, 5))) + + idata = find_MAP(model=m, progressbar=False, idata_kwargs={"include_transformed": True}) + posterior = idata.posterior.dataset.squeeze(["chain", "draw"]) + assert posterior["mu"].dims == ("group",) + assert posterior["sigma"].shape == (5,) + assert posterior["sigma_log__"].shape == (5,) + assert idata.fit.rows.values.tolist() == [ + "mu_loc", + "mu_scale_log__", + *[f"mu[{i}]" for i in range(1, 6)], + *[f"sigma_log__[{i}]" for i in range(1, 6)], + ] + + +def test_find_MAP_nonscalar_rv_without_dims(): + with pm.Model(coords={"test": ["A", "B", "C"]}) as model: + x_loc = pm.Normal("x_loc", mu=0, sigma=1, dims=["test"]) + x = pm.Normal("x", mu=x_loc, sigma=1, shape=(2, 3)) + pm.Normal("y", mu=x, sigma=1, observed=np.random.randn(10, 2, 3)) + + idata = find_MAP(model=model, progressbar=False) + assert idata.posterior["x"].shape == (1, 1, 2, 3) + assert idata.fit.rows.values.tolist() == [ + "x_loc[A]", + "x_loc[B]", + "x_loc[C]", + *[f"x[{i},{j}]" for i in range(2) for j in range(3)], + ] + + +def test_find_MAP_jitter_escapes_saddle(): + # https://github.com/pymc-devs/pymc-extras/issues/687 + with pm.Model() as m: + w = pm.Normal("w") + z = pm.Normal("z") + pm.Normal("y", mu=w * z, sigma=0.1, observed=1.0) + + stuck = map_point(model=m, jitter=False) + assert_allclose([stuck["w"], stuck["z"]], 0.0) + with pytest.warns(FutureWarning, match="will jitter its initial point"): + default = pm.find_MAP(model=m, return_inferencedata=True, progressbar=False).posterior + assert_allclose([default["w"], default["z"]], 0.0) # no jitter by default, for now + r1 = map_point(model=m, random_seed=11) + r2 = map_point(model=m, random_seed=11) + assert_allclose(r1["w"] * r1["z"], 1.0, atol=0.05) + assert_allclose(r1["w"], r2["w"]) + + +@pytest.mark.parametrize("jitter", [True, False]) +def test_find_MAP_invalid_start_raises(jitter): + with pm.Model() as m: + pm.Uniform("x", 0, 1, default_transform=None) + with pytest.raises(SamplingError, match="Initial evaluation of model at starting point failed"): + find_MAP(model=m, initvals={"x": 2.0}, jitter=jitter, progressbar=False) + + +def test_find_MAP_invalid_vars(): + with pm.Model() as m: + pm.Poisson("k", 3) + with pytest.raises(ValueError, match="no unobserved continuous variables"): + find_MAP(model=m, progressbar=False) + with pytest.raises(ValueError): + find_MAP(model=m, vars=[pt.constant(1.0)], progressbar=False) + + +def test_find_MAP_unknown_method(normal_model): + with pytest.raises(ValueError, match="Unknown method"): + find_MAP(method="gradient-descent", model=normal_model, progressbar=False) + with pytest.raises(TypeError, match="callable `method` is no longer supported"): + find_MAP(method=lambda *args, **kwargs: None, model=normal_model, progressbar=False) + + +def test_find_MAP_deprecated_defaults(normal_model): + with ( + pytest.warns(FutureWarning, match="will return an `InferenceData`"), + pytest.warns(FutureWarning, match="will jitter its initial point"), + ): + point = pm.find_MAP(model=normal_model, progressbar=False) + # the old behavior: a dict, transformed values included + assert set(point) == {"mu", "sigma", "sigma_log__"} + + +def test_find_MAP_initvals_variable_keys(normal_model): + # Variable keys must survive model freezing; maxiter=1 keeps the result near the start + kwargs = {"model": normal_model, "jitter": False, "maxiter": 1, "progressbar": False} + with pytest.warns(UserWarning, match="did not converge"): + by_var = find_MAP(initvals={normal_model["mu"]: -50.0}, **kwargs) + with pytest.warns(UserWarning, match="did not converge"): + by_name = find_MAP(initvals={"mu": -50.0}, **kwargs) + assert_allclose(by_var.posterior["mu"], by_name.posterior["mu"]) + assert by_var.posterior["mu"].item() < -10 + + +@pytest.mark.parametrize( + "legacy, new", + [ + ({"start": {"mu": 1.0}}, {"initvals": {"mu": 2.0}}), + ({"seed": 1}, {"random_seed": 2}), + ({"maxeval": 10}, {"maxiter": 20}), + ], +) +def test_find_MAP_legacy_and_new_kwargs_conflict(normal_model, legacy, new): + with pytest.raises(ValueError, match="Cannot pass both"): + find_MAP(model=normal_model, progressbar=False, **legacy, **new) + + +def test_find_MAP_legacy_kwargs(normal_model): + kwargs = {"model": normal_model, "progressbar": False, "random_seed": 1} + with pytest.warns(FutureWarning, match="`start` is deprecated"): + r1 = map_point({"mu": 1.0}, **kwargs) + with pytest.warns(FutureWarning, match="`start` is deprecated"): + r2 = map_point(start={"mu": 1.0}, **kwargs) + assert_allclose(r1["mu"], r2["mu"]) + with pytest.warns(FutureWarning, match="`seed` is deprecated"): + find_MAP(seed=1, model=normal_model, progressbar=False) + with pytest.warns(FutureWarning, match="`maxeval` is deprecated"): + find_MAP(maxeval=10, **kwargs) + with pytest.warns(FutureWarning, match="`return_raw` is deprecated"): + idata, res = find_MAP(return_raw=True, **kwargs) + with pytest.warns(FutureWarning, match="`progressbar_theme` is ignored"): + find_MAP(progressbar_theme="default", **kwargs) + with pytest.warns(FutureWarning, match="`include_transformed` is deprecated"): + legacy = find_MAP(include_transformed=True, **kwargs) + assert "sigma_log__" in legacy.posterior + with pytest.raises(ValueError, match="only via `idata_kwargs`"): + find_MAP(include_transformed=True, idata_kwargs={"include_transformed": False}, **kwargs) + # pm.sample's string progress bar options are accepted + find_MAP(**{**kwargs, "progressbar": "split+stats"}) + assert isinstance(res, OptimizeResult) + assert set(idata.posterior) == {"mu", "sigma"} + + +class TestOptimizerResultToDataset: + names = ["mu", "sigma_log__"] + + def test_basic(self): + result = OptimizeResult( + x=np.array([1.0, 2.0]), + fun=0.5, + success=True, + message="done", + jac=np.array([0.1, 0.2]), + nit=5, + custom_stat=np.array([42, 43, 44]), + status=None, + ) + ds = _optimizer_result_to_dataset(result, "BFGS", self.names) + assert isinstance(ds, xr.Dataset) + assert ds["x"].coords["variables"].values.tolist() == self.names + assert ds["jac"].dims == ("variables",) + assert ds["message"].item() == "done" and ds["method"].item() == "BFGS" + assert ds["custom_stat"].dims == ("custom_stat_dim_0",) + assert "variables_aux" not in ds.coords + + def test_lbfgs_hess_inv_kept_low_rank(self): + rng = np.random.default_rng(0) + n, m = 2, 3 + sk, yk = rng.normal(size=(m, n)), rng.normal(size=(m, n)) + yk = yk + np.sign(np.sum(sk * yk, axis=1))[:, None] * sk # keep s.y > 0 + result = OptimizeResult(x=np.ones(n), hess_inv=LbfgsInvHessProduct(sk, yk)) + ds = _optimizer_result_to_dataset(result, "L-BFGS-B", self.names) + assert "hess_inv" not in ds + for key, pairs in (("hess_inv_sk", sk), ("hess_inv_yk", yk)): + assert ds[key].dims == ("lbfgs_corrections", "variables") + assert ds[key].coords["variables"].values.tolist() == self.names + assert_allclose(ds[key].values, pairs) + + def test_basinhopping_nested_result(self): + result = OptimizeResult( + x=np.ones(2), + lowest_optimization_result=OptimizeResult(x=np.zeros(2), hess_inv=3 * np.eye(2)), + ) + ds = _optimizer_result_to_dataset(result, "basinhopping", self.names) + assert_allclose(ds["hess_inv"].values, 3 * np.eye(2)) + assert_allclose(ds["x"].values, 1.0) + assert "lowest_optimization_result" not in ds + + def test_ragged_and_mismatched_fields(self): + result = OptimizeResult( + x=np.ones(2), + method="tr_interior_point", # trust-constr's sub-method must not overwrite ours + jac=[], # trust-constr's (empty) constraint Jacobians are not per-parameter + final_simplex=(np.zeros((3, 2)), np.zeros(3)), # nelder-mead + ) + ds = _optimizer_result_to_dataset(result, "trust-constr", self.names) + assert ds["method"].item() == "trust-constr" + assert ds["jac"].dims == ("jac_dim_0",) + assert ds["final_simplex_0"].dims == ("final_simplex_0_dim_0", "variables") + assert ds["final_simplex_1"].dims == ("final_simplex_1_dim_0",)