TensorFlow新手:如何在TFP中基于自定义函数实现MCMC采样
Hey there! Let's tackle your TensorFlow Probability (TFP) questions one by one—this stuff can feel a bit daunting at first, but once you get the hang of the core patterns, it's really flexible. Here's how to handle each of your asks:
If your desired distribution can't be built by combining TFP's pre-made distributions (like mixtures or transformations), you'll need to subclass tfp.distributions.Distribution and implement the core methods TFP relies on: _sample_n (for generating samples) and log_prob (for computing log densities).
Let's use a custom bimodal distribution as an example—samples come from either a N(0,1) or N(5,1) with equal probability:
import tensorflow as tf import tensorflow_probability as tfp tfd = tfp.distributions class CustomBimodal(tfd.Distribution): def __init__(self, name="CustomBimodal"): super().__init__( dtype=tf.float32, reparameterization_type=tfd.FULLY_REPARAMETERIZED, validate_args=False, allow_nan_stats=True, name=name ) def _sample_n(self, n, seed=None): # Flip a "coin" for each sample to pick the distribution pick_first = tf.random.uniform(shape=(n,)) < 0.5 # Sample from both normals samples_a = tfd.Normal(loc=0., scale=1.).sample(n, seed=seed) samples_b = tfd.Normal(loc=5., scale=1.).sample(n, seed=seed) # Combine samples based on our coin flip return tf.where(pick_first, samples_a, samples_b) def _log_prob(self, value): # Calculate log probability of the mixture: log(0.5*p_a + 0.5*p_b) log_prob_a = tfd.Normal(loc=0., scale=1.).log_prob(value) log_prob_b = tfd.Normal(loc=5., scale=1.).log_prob(value) return tf.math.log(tf.exp(log_prob_a)*0.5 + tf.exp(log_prob_b)*0.5) # Test it out dist = CustomBimodal() samples = dist.sample(1000) print("First 5 samples:", samples[:5].numpy())
Setting reparameterization_type correctly matters if you plan to use this distribution with gradient-based methods (like MCMC or VI)—FULLY_REPARAMETERIZED means samples can be differentiated through, which is ideal.
target_log_prob Good news: TFP's MCMC samplers don't require you to wrap your target density in a Distribution class. All you need is a function that takes a tensor of parameters and returns their log probability under your target distribution.
The only critical rule: your function must be differentiable (since samplers like HMC/NUTS use gradients to propose new samples). As long as you use TensorFlow operations (avoid raw Python loops or non-differentiable functions), this works automatically.
Example of a custom 2D Gaussian mixture target log probability:
def custom_target_log_prob(params): # params is a tensor shaped [batch_size, 2] (2D points) x, y = tf.unstack(params, axis=-1) # Log probabilities for two 2D Gaussians log_prob_1 = tfd.MultivariateNormalDiag(loc=[0., 0.], scale_diag=[1., 1.]).log_prob(params) log_prob_2 = tfd.MultivariateNormalDiag(loc=[3., 3.], scale_diag=[0.5, 0.5]).log_prob(params) # 30% weight on first Gaussian, 70% on second return tf.math.log(0.3 * tf.exp(log_prob_1) + 0.7 * tf.exp(log_prob_2))
You can pass this function directly to any TFP MCMC kernel.
Let's put this into practice with the No-U-Turn Sampler (NUTS)—a robust choice for complex target densities. Here's a complete workflow:
# Set up sampling parameters num_chains = 4 # Use multiple chains to check convergence initial_state = tf.random.normal(shape=(num_chains, 2)) # Random starting points for each chain # Initialize NUTS sampler nuts_kernel = tfp.mcmc.NoUTurnSampler( target_log_prob_fn=custom_target_log_prob, step_size=0.1 # Adjust this if acceptance rates are too low/high ) # Run the sampling chain num_burnin_steps = 1000 # Warm-up phase to reach target distribution num_samples = 2000 # Samples to collect after burn-in samples, kernel_results = tfp.mcmc.sample_chain( num_results=num_samples, num_burnin_steps=num_burnin_steps, current_state=initial_state, kernel=nuts_kernel, trace_fn=lambda current_state, kr: kr # Capture sampler diagnostics ) # Process results flattened_samples = tf.reshape(samples, (-1, 2)) # Flatten all chains into one tensor # Check convergence with R-hat: values near 1 mean chains mixed well r_hat = tfp.mcmc.potential_scale_reduction(kernel_results.accepted_results.target_log_prob) print(f"R-hat convergence metric: {r_hat.numpy()}") # Calculate posterior mean posterior_mean = tf.reduce_mean(flattened_samples, axis=0) print(f"Posterior mean of 2D parameters: {posterior_mean.numpy()}")
Quick Tips:
- If your target function is non-differentiable, use a gradient-free sampler like
tfp.mcmc.RandomWalkMetropolisinstead. - Keep an eye on acceptance rates (from
kernel_results.accepted_results.accepted)—aim for ~60-80% for NUTS. If it's too low, reduce the step size; too high, increase it.
内容的提问来源于stack exchange,提问作者Jonathan Fischoff

