You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.13 09:04:09