Repository navigation
Modernize find_MAP #8429
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Modernize find_MAP #8429
Changes from all commits
dbf192a
c949146
177cee4
22f3d85
c2eb303
8077e9f
53a8822
e12b90f
3f1333a
bbd4669
7574860
9c586e1
8719f71
ee8a8c6
7d9d7ad
8ae3c07
1a21802
80e3270
e0d22d5
8f87b56
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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,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. | ||
|
|
@@ -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( | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. same
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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], | ||
|
|
@@ -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( | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. The
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Done in 38a4586: |
||
| 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]) | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
jitter=False.advi_mapwas unjittered before, and onlyjitter+mapshould jitter.There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Done in 38a4586.