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

如何在Stable Baselines3的PPO中集成Dirichlet分布?

解决SB3中PPO集成Dirichlet分布动作选择的问题

你的自定义策略类没生效,是因为Stable Baselines3(SB3)的ActorCriticPolicy并不调用你重写的action方法,它内部有一套标准的分布构建和动作采样流程。要替换成Dirichlet分布,需要重写策略类中与动作分布相关的核心方法,以下是修正方案:

核心问题分析

SB3的策略类通过_get_action_distribution生成动作分布,再通过_sample_action完成采样,同时需要计算动作的对数概率、熵等用于PPO的损失计算。你直接重写action方法完全绕开了这套流程,所以自定义逻辑不会被触发。

修正后的完整代码

1. 玩具环境(保留你的实现)

import torch
import torch.nn as nn
import gymnasium as gym
from stable_baselines3 import PPO
from stable_baselines3.common.policies import ActorCriticPolicy
from stable_baselines3.common.distributions import Distribution
import numpy as np

class ToyDirichletEnv(gym.Env):
    def __init__(self):
        super(ToyDirichletEnv, self).__init__()
        self.action_space = gym.spaces.Box(low=0, high=1, shape=(3,), dtype=np.float32)
        self.observation_space = gym.spaces.Box(low=np.array([0, 0, 0, 0]), high=np.array([10]*4), dtype=np.float32)
        self.position = 5

    def reset(self, seed=None):
        super().reset(seed=seed)
        self.position = 5
        return np.array([self.position]*4).astype(np.float32), {}

    def step(self, action):
        print('sum(action)', np.sum(action))
        reward = sum([el**2 for el in action])
        self.position += 1
        done = self.position > 15
        return np.array([self.position]*4).astype(np.float32), reward, done, done, {}

2. 自定义Dirichlet分布类

实现符合SB3接口的Dirichlet分布,处理采样、对数概率计算等核心逻辑:

class DirichletDistribution(Distribution):
    def __init__(self, alpha: torch.Tensor):
        super().__init__()
        # 确保alpha为正,避免分布报错
        self.alpha = torch.clamp(alpha, min=1e-6)
        self.distribution = torch.distributions.Dirichlet(self.alpha)

    def sample(self) -> torch.Tensor:
        return self.distribution.sample()

    def log_prob(self, actions: torch.Tensor) -> torch.Tensor:
        return self.distribution.log_prob(actions)

    def entropy(self) -> torch.Tensor:
        return self.distribution.entropy()

    def mode(self) -> torch.Tensor:
        # 返回Dirichlet分布的众数(alpha>1时有效)
        alpha_sum = self.alpha.sum(dim=-1, keepdim=True)
        return (self.alpha - 1) / (alpha_sum - self.alpha.shape[-1])

3. 集成Dirichlet的自定义策略类

重写_get_action_dist_from_latent方法,将网络输出转换为Dirichlet分布的α参数:

class DirichletActorCriticPolicy(ActorCriticPolicy):
    def __init__(self, *args, alpha_scale=0.1, **kwargs):
        super().__init__(*args, **kwargs)
        # alpha_scale控制分布集中度,值越小动作越分散
        self.alpha_scale = alpha_scale

    def _get_action_dist_from_latent(self, latent_pi: torch.Tensor, latent_vf: torch.Tensor) -> tuple[Distribution, torch.Tensor]:
        # 用softplus确保输出为正,再乘以缩放系数得到alpha
        action_logits = self.action_net(latent_pi)
        alpha = self.alpha_scale * torch.nn.functional.softplus(action_logits) + 1e-6
        # 构建Dirichlet分布
        dist = DirichletDistribution(alpha)
        # 计算价值函数输出
        value = self.value_net(latent_vf)
        return dist, value

4. 训练与验证

if __name__ == "__main__":
    env = ToyDirichletEnv()
    # 用自定义策略初始化PPO,可调整alpha_scale参数
    model = PPO(
        DirichletActorCriticPolicy,
        env,
        verbose=1,
        policy_kwargs={"alpha_scale": 0.5}
    )
    print(model.policy)
    model.learn(total_timesteps=200)

关键说明

  • 标准扩展点:通过_get_action_dist_from_latent替换动作分布,这是SB3策略类的官方扩展方式,能完美兼容PPO的训练流程。
  • 数值稳定性:用softplus和clamp确保Dirichlet分布的α参数始终为正,避免训练中出现数值错误。
  • 动作特性:Dirichlet分布采样的动作天然满足元素和为1,环境中的打印输出会直接验证这一点。
  • PPO兼容性:自定义分布实现了sample、log_prob、entropy方法,确保PPO的策略梯度、熵正则化等损失计算正常运行。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 19:57:05