tensorflow / tensorflow/probability
Transformed RandomWalkMetropolis Does not Yield Constrained Samples when using a JointDistributionNamed model
Nobody has claimed this yet.
- Dominant language
- Jupyter Notebook
- Stars
- 4.4k
- Forks
- 1.1k
- PR merge metrics
- No merged PRs in 30d
Description
In TF1 + edward2.RandomVariable, I had implemented a hierarchical model that took about a week to fit, which made experimentation painfully inefficient.
I have since re-implemented the model using TF2 + tfp.distributions.JointDistributionNamed in order to decouple prior and likelihood log_prob calculation, and data-parallelize the likelihood log_prob calculation. Transforming bijectors are used to constrain samples to be negative or positive as needed, just like with the TF1+ed2 implementation.
I am happy with the speed, but find the parameter estimation fails to converge because of NaN level-2 samples (when validate_args=False), or the following error (when validate_args=True):
Traceback (most recent call last):
File "w99_test.py", line 43, in <module>
parallel_iterations=1)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/mcmc/sample.py", line 359, in sample_chain
parallel_iterations=parallel_iterations)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/mcmc/internal/util.py", line 394, in trace_scan
parallel_iterations=parallel_iterations)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow/python/util/deprecation.py", line 574, in new_func
return func(*args, **kwargs)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow/python/ops/control_flow_ops.py", line 2491, in while_loop_v2
return_same_structure=True)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow/python/ops/control_flow_ops.py", line 2727, in while_loop
loop_vars = body(*loop_vars)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/mcmc/internal/util.py", line 383, in _body
state = loop_fn(state, elems_array.read(i))
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/mcmc/sample.py", line 343, in _trace_scan_fn
parallel_iterations=parallel_iterations)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/mcmc/internal/util.py", line 316, in smart_for_loop
parallel_iterations=parallel_iterations
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow/python/util/deprecation.py", line 574, in new_func
return func(*args, **kwargs)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow/python/ops/control_flow_ops.py", line 2491, in while_loop_v2
return_same_structure=True)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow/python/ops/control_flow_ops.py", line 2727, in while_loop
loop_vars = body(*loop_vars)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/mcmc/internal/util.py", line 314, in <lambda>
body=lambda i, *args: [i + 1] + list(body_fn(*args)),
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/mcmc/transformed_kernel.py", line 323, in one_step
previous_kernel_results.inner_results)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/mcmc/random_walk_metropolis.py", line 421, in one_step
return self._impl.one_step(current_state, previous_kernel_results)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/mcmc/metropolis_hastings.py", line 193, in one_step
previous_kernel_results.accepted_results)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/mcmc/random_walk_metropolis.py", line 499, in one_step
next_target_log_prob = self.target_log_prob_fn(*next_state_parts) # pylint: disable=not-callable
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/mcmc/transformed_kernel.py", line 102, in transformed_log_prob_fn
tlp = log_prob_fn(*fn(state_parts))
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/distributions/joint_distribution.py", line 445, in log_prob
return self._call_log_prob(value, **unmatched_kwargs)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/distributions/distribution.py", line 944, in _call_log_prob
return self._log_prob(value, **kwargs)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/distributions/joint_distribution.py", line 387, in _log_prob
xs = self._map_measure_over_dists('log_prob', value)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/distributions/joint_distribution.py", line 403, in _map_measure_over_dists
ds, xs = self._call_flat_sample_distributions(value=value, seed=42)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/distributions/joint_distribution.py", line 415, in _call_flat_sample_distributions
ds, xs = self._flat_sample_distributions(sample_shape, seed, value)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/distributions/joint_distribution_sequential.py", line 259, in _flat_sample_distributions
ds.append(dist_fn(*xs[:i])) # Chain rule of probability.
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/distributions/joint_distribution_named.py", line 264, in _fn
return dist_fn(**kwargs)
File "w99_test.py", line 16, in <lambda>
loc=prior_mean, scale=prior_stddev, name='prior', validate_args=True)
File "<decorator-gen-120>", line 2, in __init__
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/distributions/distribution.py", line 332, in wrapped_init
default_init(self_, *args, **kwargs)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/distributions/normal.py", line 147, in __init__
name=name)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/distributions/distribution.py", line 550, in __init__
d for d in self._parameter_control_dependencies(is_init=True)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow_probability/python/distributions/normal.py", line 257, in _parameter_control_dependencies
self.scale, message='Argument `scale` must be positive.'))
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow/python/ops/check_ops.py", line 487, in assert_positive_v2
return assert_positive(x=x, summarize=summarize, message=message, name=name)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow/python/ops/check_ops.py", line 506, in assert_positive
return assert_less(zero, x, data=data, summarize=summarize)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow/python/ops/check_ops.py", line 880, in assert_less
summarize, message, name)
File "/home/uadmin/anaconda3/envs/tf-2.1/lib/python3.7/site-packages/tensorflow/python/ops/check_ops.py", line 355, in _binary_assert
message=('\n'.join(_pretty_print(d, summarize) for d in data)))
tensorflow.python.framework.errors_impl.InvalidArgumentError: Argument `scale` must be positive.
Condition x > 0 did not hold element-wise:
x (shape=(1,) dtype=float32) =
['-1.0169721']
I have gotten the same results whether using a TransformedTransitionKernel + ordinary Distributions, or an ordinary kernel with TransformedDistributions.
This error can be reproduced using the following simplified example (note that even the HalfNormal is yielding negative samples, which makes log_prob evaluation of HalfNormal samples fail):
from collections import OrderedDict
from tensorflow_probability import mcmc
from tensorflow_probability import bijectors as tfb
from tensorflow_probability import distributions as tfd
def trace_fn(_, results):
return {
'target_log_prob': results.inner_results.accepted_results.target_log_prob,
'log_accept_ratio': results.inner_results.log_accept_ratio,
'is_accepted': results.inner_results.is_accepted
}
num_chains = 1
num_results = 10
num_burnin_steps = 10
num_steps_between_results = 2
model_rvs = OrderedDict({
'prior_mean': tfd.Normal(loc=0., scale=1., name='prior_mean'),
'prior_stddev': tfd.HalfNormal(scale=1., name='prior_stddev'),
'prior': lambda prior_mean, prior_stddev: tfd.Normal(
loc=prior_mean, scale=prior_stddev, name='prior', validate_args=True)
})
model = tfd.JointDistributionNamed(model_rvs)
print('model: ', model.resolve_graph())
sample = [s.numpy() for s in model.sample(num_chains).values()]
print('sample: ', sample)
transforming_bijector = [tfb.Softplus(), tfb.Identity(), tfb.Softplus()]
inverted_sample = [b.inverse(s) for b, s in zip(transforming_bijector, sample)]
print('inverted_sample: ', inverted_sample)
states, trace = mcmc.sample_chain(
num_results=num_results,
num_burnin_steps=num_burnin_steps,
num_steps_between_results=num_steps_between_results,
current_state=sample,
kernel=mcmc.TransformedTransitionKernel(
inner_kernel=mcmc.RandomWalkMetropolis(
target_log_prob_fn=model.log_prob,
new_state_fn=mcmc.random_walk_uniform_fn(scale=1.5),
seed=list(range(num_chains))),
bijector=transforming_bijector),
trace_fn=trace_fn,
parallel_iterations=1)
print('states:\n\n', states)
print('trace:\n\n', trace)
I have observed the same behavior when using equivalent edward2 models.
What is going on here?
My current environment:
numpy - 1.18.3
tensorflow_probability - 0.10.1
tensorflow - 2.2.0
os - Ubuntu 18.04 LTS
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start by running the simplified example with TensorFlow Probability 0.10.1 and TensorFlow 2.2.0. Inspect the interaction between JointDistributionNamed, TransformedTransitionKernel, RandomWalkMetropolis, and the Softplus bijectors, especially the reported negative HalfNormal samples and invalid scale. Done means identifying why the constraints are lost and documenting or fixing the behavior with a regression test.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- machine-learning
- Issue type
- Bug
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Activity status
- Stale
- Clarity
- Needs clarification
- Newbie friendliness
- 35/100