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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 12:36:17