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

Ray[RLlib]使用TorchDeterministic自定义动作分布报属性错误如何解决

错误原因

你自定义的TorchDeterministic类继承了TorchDistributionWrapper,但没有初始化父类要求的dist属性,而父类默认的logp()方法会调用self.dist.log_prob(actions)计算动作的对数概率,因此触发属性不存在的报错。

解决方法

你可以选择以下任意一种方案修改你的自定义动作分布类:

方案1:重写logp方法(轻量实现)

直接在现有类中添加logp方法的重写实现,跳过对dist属性的调用:

import torch
from ray.rllib.models.torch.torch_action_dist import TorchDistributionWrapper
from ray.rllib.models.action_dist import ActionDistribution
from ray.rllib.utils.annotations import override
from ray.rllib.utils.typing import TensorType, ModelConfigDict, ModelV2
import gym
import numpy as np
from typing import Union, Optional

class TorchDeterministic(TorchDistributionWrapper):
    """Action distribution that returns the input values directly.
    This is similar to DiagGaussian with standard deviation zero (thus only
    requiring the "mean" values as NN output).
    """
    @override(TorchDistributionWrapper)
    def __init__(self, inputs: TensorType, model: Optional[ModelV2] = None):
        super().__init__(inputs, model)

    @override(ActionDistribution)
    def deterministic_sample(self) -> TensorType:
        return self.inputs

    @override(TorchDistributionWrapper)
    def sampled_action_logp(self) -> TensorType:
        return torch.zeros((self.inputs.size()[0], ), dtype=torch.float32, device=self.inputs.device)

    # 新增重写的logp方法
    @override(TorchDistributionWrapper)
    def logp(self, actions: TensorType) -> TensorType:
        # 确定性分布下,仅当输入动作和模型输出完全一致时对数概率为0,否则为负无穷
        # 如果你的训练逻辑不需要严格校验动作匹配,也可以直接返回和sampled_action_logp一致的全0值
        match = torch.all(torch.isclose(actions, self.inputs), dim=-1)
        return torch.where(
            match,
            torch.zeros(actions.shape[0], dtype=torch.float32, device=actions.device),
            torch.full((actions.shape[0],), -float("inf"), dtype=torch.float32, device=actions.device)
        )

    @override(TorchDistributionWrapper)
    def sample(self) -> TensorType:
        return self.deterministic_sample()

    @staticmethod
    @override(ActionDistribution)
    def required_model_output_shape(
            action_space: gym.Space,
            model_config: ModelConfigDict) -> Union[int, np.ndarray]:
        return np.prod(action_space.shape)

方案2:初始化dist属性(符合RLlib原生设计规范)

在初始化方法中直接构造PyTorch原生的确定性分布赋值给dist属性,后续父类的logp、entropy等方法都可以直接复用:

import torch
from ray.rllib.models.torch.torch_action_dist import TorchDistributionWrapper
from ray.rllib.models.action_dist import ActionDistribution
from ray.rllib.utils.annotations import override
from ray.rllib.utils.typing import TensorType, ModelConfigDict, ModelV2
import gym
import numpy as np
from typing import Union, Optional

class TorchDeterministic(TorchDistributionWrapper):
    """Action distribution that returns the input values directly.
    This is similar to DiagGaussian with standard deviation zero (thus only
    requiring the "mean" values as NN output).
    """
    @override(TorchDistributionWrapper)
    def __init__(self, inputs: TensorType, model: Optional[ModelV2] = None):
        super().__init__(inputs, model)
        # 初始化确定性分布,用Independent包装处理多维动作空间的维度对齐问题
        self.dist = torch.distributions.Independent(
            torch.distributions.Deterministic(loc=self.inputs),
            reinterpreted_batch_ndims=len(self.inputs.shape[1:])
        )

    @override(ActionDistribution)
    def deterministic_sample(self) -> TensorType:
        return self.inputs

    @override(TorchDistributionWrapper)
    def sampled_action_logp(self) -> TensorType:
        return torch.zeros((self.inputs.size()[0], ), dtype=torch.float32, device=self.inputs.device)

    @override(TorchDistributionWrapper)
    def sample(self) -> TensorType:
        return self.deterministic_sample()

    @staticmethod
    @override(ActionDistribution)
    def required_model_output_shape(
            action_space: gym.Space,
            model_config: ModelConfigDict) -> Union[int, np.ndarray]:
        return np.prod(action_space.shape)

注意:两种方案都补充了张量device的适配,避免CPU/GPU张量不匹配的潜在问题。


内容的提问来源于stack exchange,提问作者Paketto

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 17:54:08