tensorflow / tensorflow/probability

Transformed RandomWalkMetropolis Does not Yield Constrained Samples when using a JointDistributionNamed model

Open
#1,020 3 comments 0 reactions 0 assignees View on GitHub

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

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. 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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.