gemlib 0.9.2__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- gemlib/__init__.py +9 -0
- gemlib/distributions/__init__.py +26 -0
- gemlib/distributions/brownian.py +141 -0
- gemlib/distributions/categorical2.py +33 -0
- gemlib/distributions/continuous_markov.py +371 -0
- gemlib/distributions/continuous_time_state_transition_model.py +185 -0
- gemlib/distributions/continuous_time_state_transition_model_test.py +289 -0
- gemlib/distributions/discrete_markov.py +279 -0
- gemlib/distributions/discrete_rejection_sampling.py +149 -0
- gemlib/distributions/discrete_time_state_transition_model.py +324 -0
- gemlib/distributions/discrete_time_state_transition_model_examples.py +453 -0
- gemlib/distributions/discrete_time_state_transition_model_test.py +336 -0
- gemlib/distributions/experimental/__init__.py +7 -0
- gemlib/distributions/experimental/discrete_approx_cont_state_transition_model.py +194 -0
- gemlib/distributions/experimental/state_transition_marginal_model.py +352 -0
- gemlib/distributions/hypergeometric.py +144 -0
- gemlib/distributions/hypergeometric_sampler.py +103 -0
- gemlib/distributions/hypergeometric_test.py +46 -0
- gemlib/distributions/kcategorical.py +113 -0
- gemlib/distributions/kcategorical_test.py +52 -0
- gemlib/distributions/uniform_integer.py +169 -0
- gemlib/distributions/uniform_integer_test.py +55 -0
- gemlib/mcmc/__init__.py +23 -0
- gemlib/mcmc/adaptive_random_walk_metropolis.py +859 -0
- gemlib/mcmc/adaptive_random_walk_metropolis_test.py +144 -0
- gemlib/mcmc/bb_fixture.pkl +0 -0
- gemlib/mcmc/brownian_bridge_kernel.py +291 -0
- gemlib/mcmc/brownian_bridge_kernel_test.py +164 -0
- gemlib/mcmc/chain_binomial_rippler.py +524 -0
- gemlib/mcmc/chain_binomial_rippler_test.py +120 -0
- gemlib/mcmc/compound_kernel.py +156 -0
- gemlib/mcmc/conftest.py +4 -0
- gemlib/mcmc/damped_chain_binomial_rippler.py +840 -0
- gemlib/mcmc/discrete_time_state_transition_model/__init__.py +21 -0
- gemlib/mcmc/discrete_time_state_transition_model/event_time_proposal.py +275 -0
- gemlib/mcmc/discrete_time_state_transition_model/fixtures.py +56 -0
- gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_mh.py +258 -0
- gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_mh_test.py +108 -0
- gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_proposal.py +188 -0
- gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_proposal_test.py +33 -0
- gemlib/mcmc/discrete_time_state_transition_model/move_events.py +239 -0
- gemlib/mcmc/discrete_time_state_transition_model/move_events_test.py +63 -0
- gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_mh.py +254 -0
- gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_mh_test.py +28 -0
- gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_proposal.py +150 -0
- gemlib/mcmc/discrete_time_state_transition_model/util.py +9 -0
- gemlib/mcmc/experimental/__init__.py +0 -0
- gemlib/mcmc/experimental/composable_kernel.py +270 -0
- gemlib/mcmc/experimental/discrete_time_state_transition_model/__init__.py +1 -0
- gemlib/mcmc/experimental/discrete_time_state_transition_model/left_censored_events_mh.py +149 -0
- gemlib/mcmc/experimental/discrete_time_state_transition_model/move_events.py +71 -0
- gemlib/mcmc/experimental/discrete_time_state_transition_model/move_events_test.py +48 -0
- gemlib/mcmc/experimental/discrete_time_state_transition_model/right_censored_events_mh.py +89 -0
- gemlib/mcmc/experimental/discrete_time_state_transition_model/right_censored_events_mh_test.py +28 -0
- gemlib/mcmc/experimental/hmc.py +111 -0
- gemlib/mcmc/experimental/hmc_test.py +61 -0
- gemlib/mcmc/experimental/mcmc_base.py +33 -0
- gemlib/mcmc/experimental/mcmc_sampler.py +94 -0
- gemlib/mcmc/experimental/mcmc_sampler_test.py +44 -0
- gemlib/mcmc/experimental/multi_scan.py +53 -0
- gemlib/mcmc/experimental/multi_scan_test.py +71 -0
- gemlib/mcmc/experimental/random_walk_metropolis.py +107 -0
- gemlib/mcmc/experimental/random_walk_metropolis_test.py +214 -0
- gemlib/mcmc/experimental/test_util.py +49 -0
- gemlib/mcmc/gibbs_kernel.py +505 -0
- gemlib/mcmc/gibbs_kernel_test.py +212 -0
- gemlib/mcmc/h5_posterior.py +77 -0
- gemlib/mcmc/multi_scan_kernel.py +59 -0
- gemlib/mcmc/zarr_posterior.py +132 -0
- gemlib/util.py +117 -0
- gemlib/util_test.py +75 -0
- gemlib-0.9.2.dist-info/LICENSE +21 -0
- gemlib-0.9.2.dist-info/METADATA +19 -0
- gemlib-0.9.2.dist-info/RECORD +75 -0
- gemlib-0.9.2.dist-info/WHEEL +4 -0
|
@@ -0,0 +1,859 @@
|
|
|
1
|
+
"""Adaptive Multisite Random Walk Metropolis Transition Kernel.
|
|
2
|
+
Authors: Alison Hale and Chris Jewell
|
|
3
|
+
Version: 0.0.21
|
|
4
|
+
Date: 18/12/2020
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
import collections
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
import tensorflow as tf
|
|
11
|
+
from tensorflow_probability.python.distributions.bernoulli import Bernoulli
|
|
12
|
+
from tensorflow_probability.python.distributions.mvn_tril import (
|
|
13
|
+
MultivariateNormalTriL,
|
|
14
|
+
)
|
|
15
|
+
from tensorflow_probability.python.distributions.normal import Normal
|
|
16
|
+
from tensorflow_probability.python.experimental import (
|
|
17
|
+
stats,
|
|
18
|
+
)
|
|
19
|
+
|
|
20
|
+
# ***tfp nightly experimental***
|
|
21
|
+
from tensorflow_probability.python.internal import dtype_util, unnest
|
|
22
|
+
from tensorflow_probability.python.mcmc import kernel as kernel_base
|
|
23
|
+
from tensorflow_probability.python.mcmc import (
|
|
24
|
+
metropolis_hastings,
|
|
25
|
+
random_walk_metropolis,
|
|
26
|
+
)
|
|
27
|
+
from tensorflow_probability.python.mcmc.internal import util as mcmc_util
|
|
28
|
+
|
|
29
|
+
__all__ = [
|
|
30
|
+
"AdaptiveRWMResults",
|
|
31
|
+
"rwm_extra_getter_fn",
|
|
32
|
+
"rwm_extra_setter_fn",
|
|
33
|
+
"rwm_log_accept_prob_getter_fn",
|
|
34
|
+
"random_walk_mvnorm_fn",
|
|
35
|
+
"AdaptiveRandomWalkMetropolis",
|
|
36
|
+
]
|
|
37
|
+
|
|
38
|
+
COV_SCALE_REDUCER_MIN = 0.5
|
|
39
|
+
COV_SCALE_REDUCER_MAX = 1.0
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class AdaptiveRWMResults(
|
|
43
|
+
mcmc_util.PrettyNamedTupleMixin,
|
|
44
|
+
collections.namedtuple(
|
|
45
|
+
"AdaptiveRWMResults",
|
|
46
|
+
[
|
|
47
|
+
"num_steps",
|
|
48
|
+
"covariance_scaling",
|
|
49
|
+
"covariance",
|
|
50
|
+
"running_covariance",
|
|
51
|
+
"is_adaptive",
|
|
52
|
+
],
|
|
53
|
+
),
|
|
54
|
+
):
|
|
55
|
+
"""State information for `MetropolisHastings` `extra` attribute.
|
|
56
|
+
|
|
57
|
+
Attributes
|
|
58
|
+
----------
|
|
59
|
+
num_steps: Python integer representing the current number of
|
|
60
|
+
`MetropolisHastings` steps.
|
|
61
|
+
covariance_scaling: Python floating point number representing a
|
|
62
|
+
the value of the prefactor which is used at each
|
|
63
|
+
`covariance_update_inverval`. The `covariance_scaling`
|
|
64
|
+
is tuned during the evolution of the MCMC chain. Let d represent the
|
|
65
|
+
number of parameters e.g. as given by the initial_state. The float given
|
|
66
|
+
by the `covariance_scaling` divided by d is used to multiply the running
|
|
67
|
+
covariance at each `covariance_update_inverval`. The result is an
|
|
68
|
+
updated covariance matrix which is used as the proposal during the
|
|
69
|
+
current `covariance_update_inverval`.
|
|
70
|
+
Default value: 2.38**2.
|
|
71
|
+
covariance: Python `list` of `Tensor`s representing the current
|
|
72
|
+
covariance of the proposal.
|
|
73
|
+
running_covariance: running covariance state as stored by
|
|
74
|
+
the instantiation of `RunningCovariance` from `stats`.
|
|
75
|
+
is_adaptive: Python `list` of `Tensor`s representing the type of
|
|
76
|
+
proposal where for each batch 0 represents a fixed proposal and 1
|
|
77
|
+
an adaptive proposal.
|
|
78
|
+
|
|
79
|
+
"""
|
|
80
|
+
|
|
81
|
+
__slots__ = ()
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def rwm_extra_getter_fn(kernel_results):
|
|
85
|
+
"""Getter for `extra` member of `MetropolisHastings` `TransitionKernel`
|
|
86
|
+
so that it can be inspected.
|
|
87
|
+
"""
|
|
88
|
+
return unnest.get_innermost(kernel_results, "extra")
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def rwm_extra_setter_fn(
|
|
92
|
+
kernel_results,
|
|
93
|
+
num_steps,
|
|
94
|
+
covariance_scaling,
|
|
95
|
+
covariance,
|
|
96
|
+
running_covariance,
|
|
97
|
+
is_adaptive,
|
|
98
|
+
):
|
|
99
|
+
"""Setter for `extra` member of `MetropolisHastings` `TransitionKernel`
|
|
100
|
+
so that it can be adapted.
|
|
101
|
+
"""
|
|
102
|
+
return unnest.replace_innermost(
|
|
103
|
+
kernel_results,
|
|
104
|
+
extra=AdaptiveRWMResults(
|
|
105
|
+
num_steps=num_steps,
|
|
106
|
+
covariance_scaling=covariance_scaling,
|
|
107
|
+
covariance=covariance,
|
|
108
|
+
running_covariance=running_covariance,
|
|
109
|
+
is_adaptive=is_adaptive,
|
|
110
|
+
),
|
|
111
|
+
)
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
def rwm_log_accept_prob_getter_fn(kernel_results):
|
|
115
|
+
"""Getter for `log_accept_prob` member of `MetropolisHastings`
|
|
116
|
+
`TransitionKernel` so that it can be inspected.
|
|
117
|
+
"""
|
|
118
|
+
log_accept_ratio = unnest.get_innermost(kernel_results, "log_accept_ratio")
|
|
119
|
+
safe_accept_ratio = tf.where(
|
|
120
|
+
tf.math.is_finite(log_accept_ratio),
|
|
121
|
+
log_accept_ratio,
|
|
122
|
+
tf.constant(-np.inf, dtype=log_accept_ratio.dtype),
|
|
123
|
+
)
|
|
124
|
+
return tf.minimum(safe_accept_ratio, 0.0)
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
def random_walk_mvnorm_fn(
|
|
128
|
+
covariance, pu=0.95, fixed_variance=0.01, is_adaptive=1, name=None
|
|
129
|
+
):
|
|
130
|
+
"""Returns callable that adds Multivariate Normal (MVN) noise to the input.
|
|
131
|
+
|
|
132
|
+
Args:
|
|
133
|
+
----
|
|
134
|
+
covariance: Python `list` of `Tensor`s representing each covariance
|
|
135
|
+
matrix, size d x d, of the Multivariate Normal proposal. The number
|
|
136
|
+
of parameters is d.
|
|
137
|
+
pu: Python floating point number representing the bounded convergence
|
|
138
|
+
parameter. If equal to 1, then all proposals are drawn
|
|
139
|
+
from the MVN(0, `covariance`) distribution, if less than 1,
|
|
140
|
+
proposals are drawn from MVN(0, `covariance`) with probability `pu`,
|
|
141
|
+
and MVN(0, `fixed_variance`/d) otherwise.
|
|
142
|
+
Default value: 0.95.
|
|
143
|
+
fixed_variance: Python floating point number representing the variance of
|
|
144
|
+
the fixed proposal distribution of the form MVN(0, `fixed_variance`/d).
|
|
145
|
+
Default value: 0.01.
|
|
146
|
+
is_adaptive: Python list of `Tensor`s representing the type of proposal
|
|
147
|
+
where for each batch 0 represents a fixed proposal and 1 an adaptive
|
|
148
|
+
proposal.
|
|
149
|
+
Default value: 1.
|
|
150
|
+
name: Python `str` name. Given the default value of `None` the name is set
|
|
151
|
+
to `random_walk_mvnorm_fn`.
|
|
152
|
+
|
|
153
|
+
Returns:
|
|
154
|
+
-------
|
|
155
|
+
random_walk_mvnorm_fn: A callable accepting a Python `list` of `Tensor`s
|
|
156
|
+
representing the state parts of the `current_state` and an `int`
|
|
157
|
+
representing the random seed to be used to generate the proposal. The
|
|
158
|
+
callable returns two quantities. First, a `Tensor` of type integer
|
|
159
|
+
representing whether each state part was updated using the fixed
|
|
160
|
+
(value=0) or adaptive (value=1) proposal. Second, a `list` of
|
|
161
|
+
`Tensor`s, with the same-type as the input state parts, which represents
|
|
162
|
+
the proposal for the Metropolis Hastings algorithm.
|
|
163
|
+
|
|
164
|
+
"""
|
|
165
|
+
dtype = dtype_util.base_dtype(covariance[0].dtype)
|
|
166
|
+
shape = tf.stack(covariance, axis=0).shape
|
|
167
|
+
# for numerical stability ensure covariance matrix is positive semi-definite
|
|
168
|
+
covariance = covariance + 1.0e-9 * tf.eye(
|
|
169
|
+
shape[1], batch_shape=[shape[0]], dtype=dtype
|
|
170
|
+
)
|
|
171
|
+
scale_tril = tf.linalg.cholesky(covariance)
|
|
172
|
+
rv_adaptive = MultivariateNormalTriL(
|
|
173
|
+
loc=tf.zeros([shape[0], shape[1]], dtype=dtype), scale_tril=scale_tril
|
|
174
|
+
)
|
|
175
|
+
rv_fixed = Normal(
|
|
176
|
+
loc=tf.zeros([shape[0], shape[1]], dtype=dtype),
|
|
177
|
+
scale=tf.constant(fixed_variance, dtype=dtype) / shape[2],
|
|
178
|
+
)
|
|
179
|
+
|
|
180
|
+
def _fn(state_parts, seed):
|
|
181
|
+
with tf.name_scope(name or "random_walk_mvnorm_fn"):
|
|
182
|
+
|
|
183
|
+
def proposal():
|
|
184
|
+
# For parallel computation it is quicker to sample
|
|
185
|
+
# both distributions then select the result
|
|
186
|
+
rv = tf.stack(
|
|
187
|
+
[
|
|
188
|
+
rv_fixed.sample(seed=seed),
|
|
189
|
+
rv_adaptive.sample(seed=seed),
|
|
190
|
+
],
|
|
191
|
+
axis=1,
|
|
192
|
+
)
|
|
193
|
+
return tf.squeeze(
|
|
194
|
+
tf.gather(rv, is_adaptive, axis=1, batch_dims=1), axis=1
|
|
195
|
+
)
|
|
196
|
+
|
|
197
|
+
proposal_parts = tf.unstack(proposal())
|
|
198
|
+
new_state_parts = [
|
|
199
|
+
proposal_part + state_part
|
|
200
|
+
for proposal_part, state_part in zip(
|
|
201
|
+
proposal_parts, state_parts
|
|
202
|
+
)
|
|
203
|
+
]
|
|
204
|
+
return new_state_parts
|
|
205
|
+
|
|
206
|
+
return _fn
|
|
207
|
+
|
|
208
|
+
|
|
209
|
+
class AdaptiveRandomWalkMetropolis(kernel_base.TransitionKernel):
|
|
210
|
+
"""Adaptive Multisite Random Walk Metropolis Algorithm.
|
|
211
|
+
Consider a continuous multivariate random variable X of dimension d,
|
|
212
|
+
distributed according to a probability distribution function pi(x)
|
|
213
|
+
known up to a normalising constant. The general principles are
|
|
214
|
+
outlined by Roberts and Rosenthal (2009)][1]. Specifically we follow
|
|
215
|
+
Algorithm 6 of Sherlock et al. (2010)[2], in which we update the MCMC
|
|
216
|
+
chain by proposing from a multivariate Normal random variable, adapting
|
|
217
|
+
both the variance and correlation structure of the covariance matrix.
|
|
218
|
+
|
|
219
|
+
In pseudo code the algorithm is:
|
|
220
|
+
```
|
|
221
|
+
Inputs:
|
|
222
|
+
i, j iteration indices with initial values 0
|
|
223
|
+
d number of dimensions (i.e. number of parameters)
|
|
224
|
+
N total number of steps
|
|
225
|
+
X[0] initial chain state
|
|
226
|
+
S = 0.001.eye(d) = initial covariance matrix
|
|
227
|
+
m[0] = 2.38^2/d the initial variance scalar i.e. covariance_scaling/d
|
|
228
|
+
t = 0.234 = target_accept_ratio
|
|
229
|
+
c = 0.01 = covariance_scaling_limiter
|
|
230
|
+
k = 0.7 = covariance_scaling_reducer
|
|
231
|
+
pu = 0.95
|
|
232
|
+
f = 0.01 = fixed_variance
|
|
233
|
+
covariance_burnin = 100
|
|
234
|
+
pi(.) denotes probability distribution of argument
|
|
235
|
+
|
|
236
|
+
for i = 0,...,N do
|
|
237
|
+
|
|
238
|
+
// Adapt covariance_scaling, m
|
|
239
|
+
if u[i-1] < pu then // Only adapt if adaptive part was proposed
|
|
240
|
+
Let alpha[i-1] = min(1, pi(X*)/pi(X[i-1])
|
|
241
|
+
Let z = max(sgn(alpha[i-1] − t), 0) / t - 1 // NB z = (-1, 1/t-1)
|
|
242
|
+
Update m[i] = m[i-1] . exp[ z . min(c, (i-1)^(−k)) ]
|
|
243
|
+
end
|
|
244
|
+
|
|
245
|
+
// Adapt covariance matrix, S
|
|
246
|
+
if i>covariance_burnin then
|
|
247
|
+
Update j = j + 1
|
|
248
|
+
Update S[j] = Cov(X[0,...,i])
|
|
249
|
+
end
|
|
250
|
+
|
|
251
|
+
// Propose new state
|
|
252
|
+
Draw u[i] ~ Uniform(0, 1)
|
|
253
|
+
if u[i] < pu then
|
|
254
|
+
// Adaptive part
|
|
255
|
+
Draw X* ~ MVN(X[i], m[i] . S[j])
|
|
256
|
+
else
|
|
257
|
+
// Fixed part
|
|
258
|
+
Draw X* ~ MVN(X[i], f . IdentityMatrix / d)
|
|
259
|
+
end
|
|
260
|
+
|
|
261
|
+
// Perform MH accept/reject
|
|
262
|
+
Let alpha[i] = min(1, pi(X*)/pi(X[i])
|
|
263
|
+
Draw v ~ Uniform(0, 1)
|
|
264
|
+
if v < alpha[i] then
|
|
265
|
+
Update X[i+1] = X*
|
|
266
|
+
else
|
|
267
|
+
Update X[i+1] = X[i]
|
|
268
|
+
end
|
|
269
|
+
|
|
270
|
+
Update i = i + 1
|
|
271
|
+
|
|
272
|
+
end
|
|
273
|
+
```
|
|
274
|
+
|
|
275
|
+
#### Example
|
|
276
|
+
```python
|
|
277
|
+
import numpy as np
|
|
278
|
+
import tensorflow as tf
|
|
279
|
+
import tensorflow_probability as tfp
|
|
280
|
+
|
|
281
|
+
tfd = tfp.distributions
|
|
282
|
+
tfb = tfp.bijectors
|
|
283
|
+
|
|
284
|
+
bijector = True
|
|
285
|
+
|
|
286
|
+
dtype = np.float32
|
|
287
|
+
|
|
288
|
+
# data
|
|
289
|
+
x = dtype([2.9, 4.2, 8.3, 1.9, 2.6, 1.0, 8.4, 8.6, 7.9, 4.3])
|
|
290
|
+
y = dtype([6.2, 7.8, 8.1, 2.7, 4.8, 2.4, 10.7, 9.0, 9.6, 5.7])
|
|
291
|
+
|
|
292
|
+
|
|
293
|
+
# define linear regression model
|
|
294
|
+
def Model(x):
|
|
295
|
+
def alpha():
|
|
296
|
+
return tfd.Normal(loc=dtype(0.0), scale=dtype(1000.0))
|
|
297
|
+
|
|
298
|
+
def beta():
|
|
299
|
+
return tfd.Normal(loc=dtype(0.0), scale=dtype(100.0))
|
|
300
|
+
|
|
301
|
+
def sigma():
|
|
302
|
+
return tfd.Gamma(concentration=dtype(0.1), rate=dtype(0.1))
|
|
303
|
+
|
|
304
|
+
def y(alpha, beta, sigma):
|
|
305
|
+
mu = alpha + beta * x
|
|
306
|
+
return tfd.Normal(mu, scale=sigma)
|
|
307
|
+
|
|
308
|
+
return tfd.JointDistributionNamed(
|
|
309
|
+
dict(alpha=alpha, beta=beta, sigma=sigma, y=y)
|
|
310
|
+
)
|
|
311
|
+
|
|
312
|
+
|
|
313
|
+
# target log probability of linear model
|
|
314
|
+
def log_prob(param):
|
|
315
|
+
alpha, beta, sigma = tf.unstack(param, axis=-1)
|
|
316
|
+
lp = model.log_prob(
|
|
317
|
+
{"alpha": alpha, "beta": beta, "sigma": sigma, "y": y}
|
|
318
|
+
)
|
|
319
|
+
return tf.reduce_sum(lp)
|
|
320
|
+
|
|
321
|
+
|
|
322
|
+
# posterior distribution MCMC chain
|
|
323
|
+
@tf.function
|
|
324
|
+
def posterior(
|
|
325
|
+
iterations, burnin, thinning, initial_state, initial_covariance
|
|
326
|
+
):
|
|
327
|
+
kernel = AdaptiveRandomWalkMetropolis(
|
|
328
|
+
target_log_prob_fn=log_prob,
|
|
329
|
+
initial_covariance=initial_covariance,
|
|
330
|
+
)
|
|
331
|
+
if bijector is True:
|
|
332
|
+
kernel = tfp.mcmc.TransformedTransitionKernel(
|
|
333
|
+
inner_kernel=kernel,
|
|
334
|
+
bijector=tfb.Blockwise(
|
|
335
|
+
[tfb.Identity(), tfb.Exp()], block_sizes=[2, 1]
|
|
336
|
+
),
|
|
337
|
+
)
|
|
338
|
+
return tfp.mcmc.sample_chain(
|
|
339
|
+
num_results=iterations,
|
|
340
|
+
current_state=initial_state,
|
|
341
|
+
kernel=kernel,
|
|
342
|
+
num_burnin_steps=burnin,
|
|
343
|
+
num_steps_between_results=thinning,
|
|
344
|
+
parallel_iterations=1,
|
|
345
|
+
trace_fn=lambda state, results: results,
|
|
346
|
+
)
|
|
347
|
+
|
|
348
|
+
|
|
349
|
+
# initialize model
|
|
350
|
+
model = Model(x)
|
|
351
|
+
initial_state = dtype(
|
|
352
|
+
[0.1, 0.1, 0.1]
|
|
353
|
+
) # start chain at alpha=0.1, beta=0.1, sigma=0.1
|
|
354
|
+
initial_covariance = dtype(0.001) * np.eye(
|
|
355
|
+
len(initial_state), dtype=dtype
|
|
356
|
+
)
|
|
357
|
+
|
|
358
|
+
# estimate posterior distribution
|
|
359
|
+
samples, results = posterior(
|
|
360
|
+
iterations=10000,
|
|
361
|
+
burnin=0,
|
|
362
|
+
thinning=0,
|
|
363
|
+
initial_state=initial_state,
|
|
364
|
+
initial_covariance=initial_covariance,
|
|
365
|
+
)
|
|
366
|
+
|
|
367
|
+
if bijector is True:
|
|
368
|
+
results = results.inner_results
|
|
369
|
+
print("using bijector")
|
|
370
|
+
|
|
371
|
+
tf.print(
|
|
372
|
+
"\nAcceptance probability:",
|
|
373
|
+
tf.math.reduce_mean(
|
|
374
|
+
tf.cast(results.is_accepted, dtype=tf.float32)
|
|
375
|
+
),
|
|
376
|
+
)
|
|
377
|
+
tf.print("\nalpha samples:", samples[:, 0])
|
|
378
|
+
tf.print("\nbeta samples:", samples[:, 1])
|
|
379
|
+
tf.print("\nsigma samples:", samples[:, 2])
|
|
380
|
+
```
|
|
381
|
+
|
|
382
|
+
#### References
|
|
383
|
+
|
|
384
|
+
[1]: Gareth Roberts, Jeffrey Rosenthal. Examples of Adaptive MCMC.
|
|
385
|
+
_Journal of Computational and Graphical Statistics_, 2009.
|
|
386
|
+
http://probability.ca/jeff/ftpdir/adaptex.pdf
|
|
387
|
+
|
|
388
|
+
[2]: Chris Sherlock, Paul Fearnhead, Gareth O. Roberts. The Random
|
|
389
|
+
Walk Metropolis: Linking Theory and Practice Through a Case Study.
|
|
390
|
+
_Statistical Science_, 25:172–190, 2010.
|
|
391
|
+
https://projecteuclid.org/download/pdfview_1/euclid.ss/1290175840
|
|
392
|
+
|
|
393
|
+
"""
|
|
394
|
+
|
|
395
|
+
def __init__(
|
|
396
|
+
self,
|
|
397
|
+
target_log_prob_fn,
|
|
398
|
+
initial_covariance,
|
|
399
|
+
initial_covariance_scaling=2.38**2,
|
|
400
|
+
covariance_scaling_reducer=0.7,
|
|
401
|
+
covariance_scaling_limiter=0.01,
|
|
402
|
+
covariance_burnin=100,
|
|
403
|
+
target_accept_ratio=0.234,
|
|
404
|
+
pu=0.95,
|
|
405
|
+
fixed_variance=0.01,
|
|
406
|
+
extra_getter_fn=rwm_extra_getter_fn,
|
|
407
|
+
extra_setter_fn=rwm_extra_setter_fn,
|
|
408
|
+
log_accept_prob_getter_fn=rwm_log_accept_prob_getter_fn,
|
|
409
|
+
seed=None,
|
|
410
|
+
name=None,
|
|
411
|
+
):
|
|
412
|
+
"""Initializes this transition kernel.
|
|
413
|
+
|
|
414
|
+
Args:
|
|
415
|
+
----
|
|
416
|
+
target_log_prob_fn: Python callable which takes an argument like
|
|
417
|
+
`current_state` and returns its (possibly unnormalized) log-density
|
|
418
|
+
under the target distribution.
|
|
419
|
+
initial_covariance: Python `list` of `Tensor`s each representing the
|
|
420
|
+
initial covariance matrix of the proposal.
|
|
421
|
+
The covariance matrix is tuned during the evolution of the MCMC
|
|
422
|
+
chain. Default value: `None`.
|
|
423
|
+
initial_covariance_scaling: Python floating point number representing
|
|
424
|
+
the initial value of the `covariance_scaling`. The value of
|
|
425
|
+
`covariance_scaling` is tuned during the evolution of the MCMC
|
|
426
|
+
chain. Let d represent the number of parameters e.g. as determined
|
|
427
|
+
by the `initial_covariance`. The ratio given by the
|
|
428
|
+
`covariance_scaling` divided by d is used to multiply the running
|
|
429
|
+
covariance. The covariance scaling factor multiplied by the
|
|
430
|
+
covariance matrix is used in the proposal at each step.
|
|
431
|
+
Default value: 2.38**2.
|
|
432
|
+
covariance_scaling_reducer: Python floating point number, bounded over
|
|
433
|
+
the range (0.5,1.0], representing the constant factor used during
|
|
434
|
+
the adaptation of the `covariance_scaling`.
|
|
435
|
+
Default value: 0.7.
|
|
436
|
+
covariance_scaling_limiter: Python floating point number, bounded
|
|
437
|
+
between 0.0 and 1.0, which places a limit on the maximum amount the
|
|
438
|
+
`covariance_scaling` value can be purturbed at each interaction of
|
|
439
|
+
the MCMC chain.
|
|
440
|
+
Default value: 0.01.
|
|
441
|
+
covariance_burnin: Python integer number of steps to take before
|
|
442
|
+
starting to compute the running covariance.
|
|
443
|
+
Default value: 100.
|
|
444
|
+
target_accept_ratio: Python floating point number, bounded between 0.0
|
|
445
|
+
and 1.0, representing the target acceptance probability of the
|
|
446
|
+
Metropolis–Hastings algorithm.
|
|
447
|
+
The default value of 0.234 is applicable when the number of
|
|
448
|
+
parameters is 3 or more. For the one parameter case typically the
|
|
449
|
+
'target_accept_ratio' should be set to 0.44.
|
|
450
|
+
pu: Python floating point number, bounded between 0.0 and 1.0,
|
|
451
|
+
representing the bounded convergence parameter. See
|
|
452
|
+
`random_walk_mvnorm_fn()` for details.
|
|
453
|
+
Default value: 0.95.
|
|
454
|
+
fixed_variance: Python floating point number representing the variance
|
|
455
|
+
of the fixed proposal distribution. See `random_walk_mvnorm_fn` for
|
|
456
|
+
further details.
|
|
457
|
+
Default value: 0.01.
|
|
458
|
+
extra_getter_fn: A callable with the signature
|
|
459
|
+
`(kernel_results) -> extra` where `kernel_results` are the results
|
|
460
|
+
of the `inner_kernel`, and `extra` is a nested collection of
|
|
461
|
+
`Tensor`s.
|
|
462
|
+
extra_setter_fn: A callable with the signature
|
|
463
|
+
`(kernel_results, args) -> new_kernel_results` where
|
|
464
|
+
`kernel_results` are the results of the `inner_kernel`, `args`
|
|
465
|
+
are a nested collection of `Tensor`s with the same
|
|
466
|
+
structure as returned by the `extra_getter_fn`, and
|
|
467
|
+
`new_kernel_results` are a copy of `kernel_results` with `args`
|
|
468
|
+
in the `extra` field set.
|
|
469
|
+
log_accept_prob_getter_fn: A callable with the signature
|
|
470
|
+
`(kernel_results) -> log_accept_prob` where `kernel_results` are the
|
|
471
|
+
results of the `inner_kernel`, and `log_accept_prob` is either a
|
|
472
|
+
a scalar, or has shape [num_chains].
|
|
473
|
+
seed: Python integer to seed the random number generator.
|
|
474
|
+
Default value: `None`.
|
|
475
|
+
name: Python `str` name prefixed to Ops created by this function.
|
|
476
|
+
Default value: `None`.
|
|
477
|
+
|
|
478
|
+
Returns:
|
|
479
|
+
-------
|
|
480
|
+
next_state: Tensor or list of `Tensor`s representing the state(s)
|
|
481
|
+
of the Markov chain(s) at each result step. Has same shape as
|
|
482
|
+
`current_state`.
|
|
483
|
+
kernel_results: `collections.namedtuple` of internal calculations used
|
|
484
|
+
to advance the chain.
|
|
485
|
+
|
|
486
|
+
Raises:
|
|
487
|
+
------
|
|
488
|
+
ValueError: if `initial_covariance_scaling` is less than or equal
|
|
489
|
+
to 0.0.
|
|
490
|
+
ValueError: if `covariance_scaling_reducer` is less than or equal
|
|
491
|
+
to 0.5 or greater than 1.0.
|
|
492
|
+
ValueError: if `covariance_scaling_limiter` is less than 0.0 or
|
|
493
|
+
greater than 1.0.
|
|
494
|
+
ValueError: if `covariance_burnin` is less than 0.
|
|
495
|
+
ValueError: if `target_accept_ratio` is less than 0.0 or
|
|
496
|
+
greater than 1.0.
|
|
497
|
+
ValueError: if `pu` is less than 0.0 or greater than 1.0.
|
|
498
|
+
ValueError: if `fixed_variance` is less than 0.0.
|
|
499
|
+
|
|
500
|
+
"""
|
|
501
|
+
with tf.name_scope(
|
|
502
|
+
mcmc_util.make_name(
|
|
503
|
+
name, "AdaptiveRandomWalkMetropolis", "__init__"
|
|
504
|
+
)
|
|
505
|
+
) as name:
|
|
506
|
+
if initial_covariance_scaling <= 0.0:
|
|
507
|
+
raise ValueError(
|
|
508
|
+
"`{}` must be a `float` greater than 0.0".format(
|
|
509
|
+
"initial_covariance_scaling"
|
|
510
|
+
)
|
|
511
|
+
)
|
|
512
|
+
if (
|
|
513
|
+
covariance_scaling_reducer <= COV_SCALE_REDUCER_MIN
|
|
514
|
+
or covariance_scaling_reducer > COV_SCALE_REDUCER_MAX
|
|
515
|
+
):
|
|
516
|
+
raise ValueError(
|
|
517
|
+
"`{}` must be a `float` greater than 0.5 and less than or\
|
|
518
|
+
equal to 1.0.".format("covariance_scaling_reducer")
|
|
519
|
+
)
|
|
520
|
+
if (
|
|
521
|
+
covariance_scaling_limiter < 0.0
|
|
522
|
+
or covariance_scaling_limiter > 1.0
|
|
523
|
+
):
|
|
524
|
+
raise ValueError(
|
|
525
|
+
"`{}` must be a `float` between 0.0 and 1.0.".format(
|
|
526
|
+
"covariance_scaling_limiter"
|
|
527
|
+
)
|
|
528
|
+
)
|
|
529
|
+
if covariance_burnin < 0:
|
|
530
|
+
raise ValueError(
|
|
531
|
+
"`{}` must be a `integer` greater or equal to 0.".format(
|
|
532
|
+
"covariance_burnin"
|
|
533
|
+
)
|
|
534
|
+
)
|
|
535
|
+
if target_accept_ratio <= 0.0 or target_accept_ratio > 1.0:
|
|
536
|
+
raise ValueError(
|
|
537
|
+
"`{}` must be a `float` between 0.0 and 1.0.".format(
|
|
538
|
+
"target_accept_ratio"
|
|
539
|
+
)
|
|
540
|
+
)
|
|
541
|
+
if pu < 0.0 or pu > 1.0:
|
|
542
|
+
raise ValueError(
|
|
543
|
+
"`{}` must be a `float` between 0.0 and 1.0.".format("pu")
|
|
544
|
+
)
|
|
545
|
+
if fixed_variance < 0.0:
|
|
546
|
+
raise ValueError(
|
|
547
|
+
"`{}` must be a `float` greater than 0.0.".format(
|
|
548
|
+
"fixed_variance"
|
|
549
|
+
)
|
|
550
|
+
)
|
|
551
|
+
|
|
552
|
+
if initial_covariance.shape == ():
|
|
553
|
+
initial_covariance_ = tf.reshape(initial_covariance, shape=(1, 1))
|
|
554
|
+
else:
|
|
555
|
+
initial_covariance_ = initial_covariance
|
|
556
|
+
|
|
557
|
+
if mcmc_util.is_list_like(initial_covariance_):
|
|
558
|
+
initial_covariance_parts = list(initial_covariance_)
|
|
559
|
+
else:
|
|
560
|
+
initial_covariance_parts = [initial_covariance_]
|
|
561
|
+
initial_covariance_parts = [
|
|
562
|
+
tf.convert_to_tensor(s, name="initial_covariance_")
|
|
563
|
+
for s in initial_covariance_parts
|
|
564
|
+
]
|
|
565
|
+
self._initial_covariance_matrices = tf.stack(initial_covariance_parts)
|
|
566
|
+
|
|
567
|
+
dtype = dtype_util.base_dtype(self._initial_covariance_matrices.dtype)
|
|
568
|
+
shape = self._initial_covariance_matrices.shape
|
|
569
|
+
|
|
570
|
+
self._running_covar = stats.RunningCovariance.from_shape(
|
|
571
|
+
shape=(1, shape[-1]), dtype=dtype, event_ndims=1
|
|
572
|
+
)
|
|
573
|
+
|
|
574
|
+
probs = tf.expand_dims(tf.ones([shape[0]], dtype=dtype) * pu, axis=1)
|
|
575
|
+
self._u = Bernoulli(probs=probs, dtype=tf.dtypes.int32)
|
|
576
|
+
self._initial_u = tf.zeros_like(
|
|
577
|
+
self._u.sample(seed=seed), dtype=tf.dtypes.int32
|
|
578
|
+
)
|
|
579
|
+
|
|
580
|
+
name = mcmc_util.make_name(name, "AdaptiveRandomWalkMetropolis", "")
|
|
581
|
+
|
|
582
|
+
self._parameters = {
|
|
583
|
+
"target_log_prob_fn": target_log_prob_fn,
|
|
584
|
+
"initial_covariance": initial_covariance,
|
|
585
|
+
"initial_covariance_scaling": initial_covariance_scaling,
|
|
586
|
+
"covariance_scaling_reducer": covariance_scaling_reducer,
|
|
587
|
+
"covariance_scaling_limiter": covariance_scaling_limiter,
|
|
588
|
+
"covariance_burnin": covariance_burnin,
|
|
589
|
+
"target_accept_ratio": target_accept_ratio,
|
|
590
|
+
"pu": pu,
|
|
591
|
+
"fixed_variance": fixed_variance,
|
|
592
|
+
"extra_getter_fn": extra_getter_fn,
|
|
593
|
+
"extra_setter_fn": extra_setter_fn,
|
|
594
|
+
"log_accept_prob_getter_fn": log_accept_prob_getter_fn,
|
|
595
|
+
"seed": seed,
|
|
596
|
+
"name": name,
|
|
597
|
+
}
|
|
598
|
+
self._impl = metropolis_hastings.MetropolisHastings(
|
|
599
|
+
inner_kernel=random_walk_metropolis.UncalibratedRandomWalk(
|
|
600
|
+
target_log_prob_fn=target_log_prob_fn,
|
|
601
|
+
new_state_fn=random_walk_mvnorm_fn(
|
|
602
|
+
covariance=self._initial_covariance_matrices,
|
|
603
|
+
pu=pu,
|
|
604
|
+
fixed_variance=fixed_variance,
|
|
605
|
+
is_adaptive=self._initial_u,
|
|
606
|
+
name=name,
|
|
607
|
+
),
|
|
608
|
+
name=name,
|
|
609
|
+
),
|
|
610
|
+
name=name,
|
|
611
|
+
)
|
|
612
|
+
|
|
613
|
+
@property
|
|
614
|
+
def target_log_prob_fn(self):
|
|
615
|
+
return self._parameters["target_log_prob_fn"]
|
|
616
|
+
|
|
617
|
+
@property
|
|
618
|
+
def initial_covariance(self):
|
|
619
|
+
return self._parameters["initial_covariance"]
|
|
620
|
+
|
|
621
|
+
@property
|
|
622
|
+
def initial_covariance_scaling(self):
|
|
623
|
+
return self._parameters["initial_covariance_scaling"]
|
|
624
|
+
|
|
625
|
+
@property
|
|
626
|
+
def covariance_scaling_reducer(self):
|
|
627
|
+
return self._parameters["covariance_scaling_reducer"]
|
|
628
|
+
|
|
629
|
+
@property
|
|
630
|
+
def covariance_scaling_limiter(self):
|
|
631
|
+
return self._parameters["covariance_scaling_limiter"]
|
|
632
|
+
|
|
633
|
+
@property
|
|
634
|
+
def covariance_burnin(self):
|
|
635
|
+
return self._parameters["covariance_burnin"]
|
|
636
|
+
|
|
637
|
+
@property
|
|
638
|
+
def target_accept_ratio(self):
|
|
639
|
+
return self._parameters["target_accept_ratio"]
|
|
640
|
+
|
|
641
|
+
@property
|
|
642
|
+
def pu(self):
|
|
643
|
+
return self._parameters["pu"]
|
|
644
|
+
|
|
645
|
+
@property
|
|
646
|
+
def fixed_variance(self):
|
|
647
|
+
return self._parameters["fixed_variance"]
|
|
648
|
+
|
|
649
|
+
def extra_setter_fn(
|
|
650
|
+
self,
|
|
651
|
+
kernel_results,
|
|
652
|
+
num_steps,
|
|
653
|
+
covariance_scaling,
|
|
654
|
+
covariance,
|
|
655
|
+
running_covariance,
|
|
656
|
+
is_accepted,
|
|
657
|
+
):
|
|
658
|
+
return self._parameters["extra_setter_fn"](
|
|
659
|
+
kernel_results,
|
|
660
|
+
num_steps,
|
|
661
|
+
covariance_scaling,
|
|
662
|
+
covariance,
|
|
663
|
+
running_covariance,
|
|
664
|
+
is_accepted,
|
|
665
|
+
)
|
|
666
|
+
|
|
667
|
+
def extra_getter_fn(self, kernel_results):
|
|
668
|
+
return self._parameters["extra_getter_fn"](kernel_results)
|
|
669
|
+
|
|
670
|
+
def log_accept_prob_getter_fn(self, kernel_results):
|
|
671
|
+
return self._parameters["log_accept_prob_getter_fn"](kernel_results)
|
|
672
|
+
|
|
673
|
+
@property
|
|
674
|
+
def seed(self):
|
|
675
|
+
return self._parameters["seed"]
|
|
676
|
+
|
|
677
|
+
@property
|
|
678
|
+
def name(self):
|
|
679
|
+
return self._parameters["name"]
|
|
680
|
+
|
|
681
|
+
@property
|
|
682
|
+
def parameters(self):
|
|
683
|
+
"""Return `dict` of ``__init__`` arguments and their values."""
|
|
684
|
+
return self._parameters
|
|
685
|
+
|
|
686
|
+
@property
|
|
687
|
+
def running_covar(self):
|
|
688
|
+
return self._running_covar
|
|
689
|
+
|
|
690
|
+
@property
|
|
691
|
+
def u(self):
|
|
692
|
+
return self._u
|
|
693
|
+
|
|
694
|
+
@property
|
|
695
|
+
def initial_u(self):
|
|
696
|
+
return self._initial_u
|
|
697
|
+
|
|
698
|
+
@property
|
|
699
|
+
def is_calibrated(self):
|
|
700
|
+
return True
|
|
701
|
+
|
|
702
|
+
def update_covariance_scaling(self, prev_results, num_steps):
|
|
703
|
+
previous_covar_scaling = self.extra_getter_fn(
|
|
704
|
+
prev_results
|
|
705
|
+
).covariance_scaling
|
|
706
|
+
previous_log_accept_ratio = self.log_accept_prob_getter_fn(prev_results)
|
|
707
|
+
dtype = dtype_util.base_dtype(previous_covar_scaling.dtype)
|
|
708
|
+
covariance_scaling_reducer = tf.constant(
|
|
709
|
+
self.covariance_scaling_reducer, dtype=dtype
|
|
710
|
+
)
|
|
711
|
+
covariance_scaling_limiter = tf.constant(
|
|
712
|
+
self.covariance_scaling_limiter, dtype=dtype
|
|
713
|
+
)
|
|
714
|
+
target_accept_ratio = tf.constant(self.target_accept_ratio, dtype=dtype)
|
|
715
|
+
cond = previous_log_accept_ratio - tf.math.log(target_accept_ratio)
|
|
716
|
+
multiplier = tf.math.maximum(
|
|
717
|
+
tf.math.sign(cond), tf.constant(0.0, dtype)
|
|
718
|
+
) * (tf.constant(1.0, dtype) / target_accept_ratio) - tf.constant(
|
|
719
|
+
1.0, dtype
|
|
720
|
+
)
|
|
721
|
+
delta = tf.math.minimum(
|
|
722
|
+
covariance_scaling_limiter,
|
|
723
|
+
tf.cast(num_steps, dtype=dtype) ** (-covariance_scaling_reducer),
|
|
724
|
+
)
|
|
725
|
+
return previous_covar_scaling * tf.math.exp(delta * multiplier)
|
|
726
|
+
|
|
727
|
+
def one_step(self, current_state, previous_kernel_results, seed=None):
|
|
728
|
+
with tf.name_scope(
|
|
729
|
+
mcmc_util.make_name(
|
|
730
|
+
self.name, "AdaptiveRandomWalkMetropolis", "one_step"
|
|
731
|
+
)
|
|
732
|
+
):
|
|
733
|
+
with tf.name_scope("initialize"):
|
|
734
|
+
if mcmc_util.is_list_like(current_state):
|
|
735
|
+
current_state_parts = list(current_state)
|
|
736
|
+
else:
|
|
737
|
+
current_state_parts = [current_state]
|
|
738
|
+
current_state_parts = [
|
|
739
|
+
tf.convert_to_tensor(s, name="current_state")
|
|
740
|
+
for s in current_state_parts
|
|
741
|
+
]
|
|
742
|
+
|
|
743
|
+
# Note 'covariance_scaling' and 'accum_covar' are updated every step
|
|
744
|
+
# but 'covariance' is not updated until 'num_steps' >=
|
|
745
|
+
# 'covariance_burnin'.
|
|
746
|
+
num_steps = self.extra_getter_fn(previous_kernel_results).num_steps
|
|
747
|
+
# for parallel processing efficiency use gather() vs cond()?
|
|
748
|
+
previous_is_adaptive = self.extra_getter_fn(
|
|
749
|
+
previous_kernel_results
|
|
750
|
+
).is_adaptive
|
|
751
|
+
current_covariance_scaling = tf.gather(
|
|
752
|
+
tf.stack(
|
|
753
|
+
[
|
|
754
|
+
self.extra_getter_fn(
|
|
755
|
+
previous_kernel_results
|
|
756
|
+
).covariance_scaling,
|
|
757
|
+
self.update_covariance_scaling(
|
|
758
|
+
previous_kernel_results, num_steps
|
|
759
|
+
),
|
|
760
|
+
],
|
|
761
|
+
axis=-1,
|
|
762
|
+
),
|
|
763
|
+
previous_is_adaptive,
|
|
764
|
+
batch_dims=1,
|
|
765
|
+
axis=1,
|
|
766
|
+
)
|
|
767
|
+
previous_accum_covar = self.extra_getter_fn(
|
|
768
|
+
previous_kernel_results
|
|
769
|
+
).running_covariance
|
|
770
|
+
current_accum_covar = previous_accum_covar.update(
|
|
771
|
+
new_sample=current_state_parts
|
|
772
|
+
)
|
|
773
|
+
|
|
774
|
+
previous_covariance = self.extra_getter_fn(
|
|
775
|
+
previous_kernel_results
|
|
776
|
+
).covariance
|
|
777
|
+
current_covariance = tf.gather(
|
|
778
|
+
[
|
|
779
|
+
previous_covariance,
|
|
780
|
+
current_accum_covar.covariance(ddof=1),
|
|
781
|
+
],
|
|
782
|
+
tf.cast(
|
|
783
|
+
num_steps >= self.covariance_burnin,
|
|
784
|
+
dtype=tf.dtypes.int32,
|
|
785
|
+
),
|
|
786
|
+
)
|
|
787
|
+
|
|
788
|
+
current_scaled_covariance = tf.squeeze(
|
|
789
|
+
tf.expand_dims(current_covariance_scaling, axis=1)
|
|
790
|
+
* tf.stack([current_covariance]),
|
|
791
|
+
axis=0,
|
|
792
|
+
)
|
|
793
|
+
|
|
794
|
+
current_is_adaptive = self.u.sample(seed=self.seed)
|
|
795
|
+
|
|
796
|
+
self._impl = metropolis_hastings.MetropolisHastings(
|
|
797
|
+
inner_kernel=random_walk_metropolis.UncalibratedRandomWalk(
|
|
798
|
+
target_log_prob_fn=self.target_log_prob_fn,
|
|
799
|
+
new_state_fn=random_walk_mvnorm_fn(
|
|
800
|
+
covariance=current_scaled_covariance,
|
|
801
|
+
pu=self.pu,
|
|
802
|
+
fixed_variance=self.fixed_variance,
|
|
803
|
+
is_adaptive=current_is_adaptive,
|
|
804
|
+
name=self.name,
|
|
805
|
+
),
|
|
806
|
+
name=self.name,
|
|
807
|
+
),
|
|
808
|
+
name=self.name,
|
|
809
|
+
)
|
|
810
|
+
new_state, new_inner_results = self._impl.one_step(
|
|
811
|
+
current_state, previous_kernel_results
|
|
812
|
+
)
|
|
813
|
+
new_inner_results = self.extra_setter_fn(
|
|
814
|
+
new_inner_results,
|
|
815
|
+
num_steps + 1,
|
|
816
|
+
tf.squeeze(current_covariance_scaling, axis=1),
|
|
817
|
+
current_covariance,
|
|
818
|
+
current_accum_covar,
|
|
819
|
+
current_is_adaptive,
|
|
820
|
+
)
|
|
821
|
+
return [new_state, new_inner_results]
|
|
822
|
+
|
|
823
|
+
def bootstrap_results(self, init_state):
|
|
824
|
+
"""Creates initial `state`."""
|
|
825
|
+
with tf.name_scope(
|
|
826
|
+
mcmc_util.make_name(
|
|
827
|
+
self.name, "AdaptiveRandomWalkMetropolis", "bootstrap_results"
|
|
828
|
+
)
|
|
829
|
+
):
|
|
830
|
+
if mcmc_util.is_list_like(init_state):
|
|
831
|
+
initial_state_parts = list(init_state)
|
|
832
|
+
else:
|
|
833
|
+
initial_state_parts = [init_state]
|
|
834
|
+
initial_state_parts = [
|
|
835
|
+
tf.convert_to_tensor(s, name="init_state")
|
|
836
|
+
for s in initial_state_parts
|
|
837
|
+
]
|
|
838
|
+
|
|
839
|
+
shape = tf.stack(initial_state_parts).shape
|
|
840
|
+
dtype = dtype_util.base_dtype(tf.stack(initial_state_parts).dtype)
|
|
841
|
+
|
|
842
|
+
init_covariance_scaling = tf.cast(
|
|
843
|
+
tf.repeat(
|
|
844
|
+
[self.initial_covariance_scaling],
|
|
845
|
+
repeats=[shape[0]],
|
|
846
|
+
axis=0,
|
|
847
|
+
),
|
|
848
|
+
dtype=dtype,
|
|
849
|
+
)
|
|
850
|
+
|
|
851
|
+
inner_results = self._impl.bootstrap_results(init_state)
|
|
852
|
+
return self.extra_setter_fn(
|
|
853
|
+
inner_results,
|
|
854
|
+
0,
|
|
855
|
+
init_covariance_scaling / shape[-1],
|
|
856
|
+
self._initial_covariance_matrices,
|
|
857
|
+
self._running_covar,
|
|
858
|
+
self.initial_u,
|
|
859
|
+
)
|