PyTorch贝叶斯ML中分布参数变换的计算图保留问题
Great question! The core issue here is that when you define std = torch.exp(log_std) outside your training loop, that tensor doesn't stay dynamically linked to the underlying log_std parameter—after the first backward pass, PyTorch doesn't automatically refresh std when log_std updates, breaking the computation graph and stopping gradients from flowing correctly.
Luckily, there are two clean, distribution-agnostic ways to solve this without repeating parameter transformations in every loop iteration:
1. Custom Distribution Wrapper Class (Flexible for Any Transformation)
This approach lets you encapsulate parameter transformation logic in a reusable wrapper, so you can switch distributions (like Normal → Gamma) without touching your training loop code.
import torch import torch.nn as nn import torch.distributions as dd class ParametrizedDistWrapper: def __init__(self, dist_class, param_transforms, **raw_params): self.dist_class = dist_class # Map distribution parameter names to their transformation functions self.param_transforms = param_transforms # Store raw, optimizable parameters (e.g., log_std instead of std) self.raw_params = raw_params # Collect all parameters for the optimizer self.optim_params = [p for p in raw_params.values() if isinstance(p, nn.Parameter)] def get_distribution(self): # Dynamically compute transformed parameters each time we need the distribution transformed_params = {} for param_name, raw_value in self.raw_params.items(): if param_name in self.param_transforms: transformed_params[param_name] = self.param_transforms[param_name](raw_value) else: transformed_params[param_name] = raw_value return self.dist_class(**transformed_params) # -------------------------- # Usage Example (Normal) # -------------------------- log_std = nn.Parameter(torch.Tensor([1.])) mean = nn.Parameter(torch.Tensor([1.])) # Define how to transform raw parameters to valid distribution parameters norm_wrapper = ParametrizedDistWrapper( dist_class=dd.Normal, param_transforms={"scale": torch.exp}, # Convert log_std to positive std loc=mean, scale=log_std # Pass raw log_std here ) optim = torch.optim.SGD(norm_wrapper.optim_params, lr=0.01) target_dist = dd.Normal(5., 5.) for i in range(50): optim.zero_grad() # Get the latest distribution with updated parameters current_dist = norm_wrapper.get_distribution() samples = current_dist.rsample((1000,)) # KL divergence loss cost = -(target_dist.log_prob(samples) - current_dist.log_prob(samples)).sum() cost.backward() optim.step() print(f"Iter {i}: log_std={log_std.item():.3f}, mean={mean.item():.3f}, cost={cost.item():.3f}") # -------------------------- # Switch to Gamma Distribution (No Loop Changes!) # -------------------------- log_shape = nn.Parameter(torch.Tensor([1.])) rate = nn.Parameter(torch.Tensor([1.])) gamma_wrapper = ParametrizedDistWrapper( dist_class=dd.Gamma, param_transforms={"shape": torch.exp}, # Ensure shape is positive shape=log_shape, rate=rate ) # Training loop stays identical—just swap gamma_wrapper for norm_wrapper
How It Works:
Every time you call get_distribution(), it recalculates the transformed parameters using the latest values of your raw optimizable parameters (like log_std). This ensures the computation graph is preserved, and gradients flow correctly to the raw parameters.
2. PyTorch Official Parametrization Module (Clean for Common Transforms)
For standard transformations like ensuring positivity (via exp) or clamping to a range, PyTorch's nn.utils.parametrizations module handles the dynamic parameter linking automatically—no custom classes needed.
import torch import torch.nn as nn import torch.distributions as dd from torch.nn.utils.parametrizations import exp # Use `exp` to wrap log_std: this creates a tensor that dynamically computes torch.exp(log_std) log_std = nn.Parameter(torch.Tensor([1.])) std = exp(log_std) mean = nn.Parameter(torch.Tensor([1.])) # Create the distribution once—std will auto-update with log_std! dist = dd.Normal(loc=mean, scale=std) optim = torch.optim.SGD([log_std, mean], lr=0.01) target_dist = dd.Normal(5., 5.) for i in range(50): optim.zero_grad() samples = dist.rsample((1000,)) cost = -(target_dist.log_prob(samples) - dist.log_prob(samples)).sum() cost.backward() optim.step() print(f"Iter {i}: log_std={log_std.item():.3f}, mean={mean.item():.3f}, std={std.item():.3f}, cost={cost.item():.3f}")
How It Works:
The exp(log_std) call returns a ParametrizedTensor that automatically recomputes torch.exp(log_std) every time its value is accessed. This keeps the computation graph intact, so gradients flow directly to log_std even though you're using std in the distribution.
Which to Choose?
- Custom Wrapper: Use this if you need complex, custom parameter transformations, or if you want a unified interface for multiple distribution types.
- Official Parametrization: Use this for standard transformations (exp, sigmoid, etc.)—it's more concise and leverages PyTorch's built-in tools.
内容的提问来源于stack exchange,提问作者Rabbitman14

