PyTorch中TransformedDistribution维度报错及log_prob异常的解决求助
修复PyTorch TransformedDistribution的维度匹配错误
问题代码
import torch import torch.distributions as pyd import torch.nn as nn # 修正原代码笔误:toch → torch from torch.distributions import transforms as tT from torch.distributions.transformed_distribution import TransformedDistribution import math # 补充原代码缺失导入 import torch.nn.functional as F # 补充原代码缺失导入 from torch.distributions.normal import Normal # 补充原代码缺失导入 from torch.autograd import Variable # 补充原代码缺失导入 LOG_STD_MIN = -5 LOG_STD_MAX = 0 class TanhTransform(pyd.transforms.Transform): domain = pyd.constraints.real codomain = pyd.constraints.interval(-1.0, 1.0) bijective = True sign = +1 def __init__(self, cache_size=1): self.device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu') super().__init__(cache_size=cache_size) @staticmethod def atanh(x): return 0.5 * (x.log1p() - (-x).log1p()) def __eq__(self, other): return isinstance(other, TanhTransform) def _call(self, x): return x.tanh() def _inverse(self, y): return self.atanh(y.clamp(-0.99, 0.99)) def log_abs_det_jacobian(self, x, y): return 2.0 * (math.log(2.0) - x - F.softplus(-2.0 * x)) def get_spec_means_mags(spec): means = (spec.maximum + spec.minimum) / 2.0 mags = (spec.maximum - spec.minimum) / 2.0 means = Variable(torch.tensor(means).type(torch.FloatTensor), requires_grad=False) mags = Variable(torch.tensor(mags).type(torch.FloatTensor), requires_grad=False) return means, mags class Split(torch.nn.Module): def __init__(self, module, n_parts: int, dim=1): super().__init__() self._n_parts = n_parts self._dim = dim self._module = module def forward(self, inputs): output = self._module(inputs) if output.ndim==1: result=torch.hsplit(output, self._n_parts ) else: chunk_size = output.shape[self._dim] // self._n_parts result =torch.split(output, chunk_size, dim=self._dim) return result class Network(nn.Module): def __init__( self, state, act, fc_layer_params=(), ): super(Network, self).__init__() self._act = act self._layers = nn.ModuleList() self.device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu') for hidden_size in fc_layer_params: if len(self._layers)==0: self._layers.append(nn.Linear(state.shape[0], hidden_size)) else: self._layers.append(nn.Linear(hidden_size, hidden_size)) self._layers.append(nn.ReLU()) output_layer = nn.Linear(hidden_size,self._act.shape[0] * 2) self._layers.append(output_layer) self._act_means, self._act_mags = get_spec_means_mags(self._act) # 修正变量名不匹配问题 self._action_means = self._act_means.to(self.device) self._action_mags = self._act_mags.to(self.device) def _get_outputs(self, state): h = state.to(self.device) for l in nn.Sequential(*(list(self._layers.children())[:-1])): h = l(h) self._mean_logvar_layers = Split( self._layers[-1], n_parts=2, ) mean, log_std = self._mean_logvar_layers(h) mean = mean.to(self.device) log_std = log_std.to(self.device) a_tanh_mode = torch.tanh(mean) * self._action_mags + self._action_means log_std = torch.tanh(log_std) log_std = LOG_STD_MIN + 0.5 * (LOG_STD_MAX - LOG_STD_MIN) * (log_std + 1) std = torch.exp(log_std) # 核心修正:event_dim设置为事件维度数量而非维度大小 a_distribution = TransformedDistribution( base_distribution=pyd.Normal(loc=torch.full_like(mean, 0).to(self.device), scale=torch.full_like(mean, 1).to(self.device)), transforms=tT.ComposeTransform([ tT.AffineTransform(loc=self._action_means, scale=self._action_mags, event_dim=1), TanhTransform(), tT.AffineTransform(loc=mean, scale=std, event_dim=1)])) return a_distribution, a_tanh_mode def get_log_density(self, state, action): a_dist, _ = self._get_outputs(state) log_density = a_dist.log_prob(action.to(self.device)) return log_density def __call__(self, state): a_dist, a_tanh_mode = self._get_outputs(state) a_sample = a_dist.sample() log_pi_a = a_dist.log_prob(a_sample) return a_tanh_mode, a_sample, log_pi_a
触发错误
action = self._a_network(latent_states)[1] File "/home/planner_regularizer.py", line 182, in __call__ a_dist, a_tanh_mode = self._get_outputs(state.to(device=self.device)) File "/home/planner_regularizer.py", line 159, in _get_outputs a_distribution = TransformedDistribution( File "/home/dm_control/lib/python3.8/site-packages/torch/distributions/transformed_distribution.py", line 61, in __init__ raise ValueError("base_distribution needs to have shape with size at least {}, but got {}." ValueError: base_distribution needs to have shape with size at least 6, but got torch.Size([6]).
更新说明
若移除AffineTransform中的event_dim参数,上述报错消失,但log_prob的输出尺寸为1,不符合预期。
修复方案
1. 核心错误修正:理解event_dim参数含义
AffineTransform的event_dim表示事件维度的数量,而非事件维度的大小。例如,6维动作向量属于1个事件维度(整个向量是一个独立事件),因此应设置event_dim=1,而非mean.shape[-1](即6)。
原代码错误地将维度大小传给event_dim,导致TransformedDistribution认为基础分布需要至少6个维度,但实际基础分布仅为1维张量(shape=[6]),因此触发维度不匹配错误。
2. 解决log_prob尺寸问题
移除event_dim时,默认值为0,此时变换会将整个张量视为标量事件,log_prob会对所有维度求和,最终输出尺寸为1。设置正确的event_dim=1后,log_prob会保留batch维度,每个样本对应一个独立的log概率值,符合预期。
3. 其他细节修正
- 统一变量名:原代码中
self._act_means/self._act_mags与后续使用的self._action_means/self._action_mags不匹配,需同步。 - 设备一致性:确保所有张量都移动到同一设备(CPU/GPU),避免设备不匹配报错。
- 补充缺失导入:原代码缺少
math、torch.nn.functional、Normal等模块的导入,需补充。
内容的提问来源于stack exchange,提问作者Dalek
相关产品推荐
相关产品推荐

