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

如何基于PyTorch修改离散动作空间REINFORCE算法适配连续动作空间?

适配连续动作空间的REINFORCE算法修改方案

针对你现有的离散动作空间REINFORCE代码,以下是适配连续动作空间(如Pendulum-v1)的完整修改方案,核心是将策略网络输出正态分布的参数(均值mu和标准差sigma),并基于该分布采样动作。

关键修改点

  • 策略网络输出调整:不再输出动作概率,而是输出每个动作维度的均值mu和标准差sigma,标准差需通过激活函数保证为正数
  • 采样分布替换:将离散的Categorical分布替换为连续的Normal分布
  • 动作范围裁剪:针对连续环境的动作约束(如Pendulum-v1的动作范围是[-2, 2]),对采样后的动作进行裁剪
  • 日志概率计算:基于正态分布计算动作的对数概率,用于后续损失计算

完整修改代码

import numpy as np
import torch as T
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
import gym

class PolicyNetwork(nn.Module):
    def __init__(self, lr, input_dims, action_dim):
        super(PolicyNetwork, self).__init__()
        self.fc1 = nn.Linear(*input_dims, 128)
        self.fc2 = nn.Linear(128, 128)
        # 输出维度为2*action_dim:前action_dim个是mu,后action_dim个是sigma的原始输出
        self.fc3 = nn.Linear(128, 2 * action_dim)
        self.optimizer = optim.Adam(self.parameters(), lr=lr)

        self.device = T.device('cuda:0' if T.cuda.is_available() else 'cpu')
        self.to(self.device)

    def forward(self, state):
        x = F.relu(self.fc1(state))
        x = F.relu(self.fc2(x))
        # 获取mu和原始sigma输出
        mu_sigma = self.fc3(x)
        mu = mu_sigma[:, :self.fc3.out_features//2]
        # 使用softplus保证sigma为正,避免使用exp导致数值不稳定
        sigma = F.softplus(mu_sigma[:, self.fc3.out_features//2:]) + 1e-6  # 加小值防止sigma为0
        return mu, sigma

class PolicyGradientAgent():
    def __init__(self, lr, input_dims, gamma=0.99, action_dim=1):
        self.gamma = gamma
        self.lr = lr
        self.reward_memory = []
        self.log_prob_memory = []  # 改为存储log_prob,名称更贴合实际用途

        self.policy = PolicyNetwork(self.lr, input_dims, action_dim)

    def choose_action(self, observation):
        state = T.Tensor([observation]).to(self.policy.device)
        mu, sigma = self.policy.forward(state)
        # 创建正态分布
        dist = T.distributions.Normal(mu, sigma)
        # 采样动作
        action = dist.sample()
        # 计算动作的对数概率
        log_prob = dist.log_prob(action)
        self.log_prob_memory.append(log_prob)

        # 针对Pendulum-v1裁剪动作到[-2, 2]
        action_clipped = T.clamp(action, -2, 2)
        return action_clipped.item()

    def store_rewards(self, reward):
        self.reward_memory.append(reward)

    def learn(self):
        self.policy.optimizer.zero_grad()

        # 计算折扣回报G_t,逻辑与原代码一致
        G = np.zeros_like(self.reward_memory, dtype=np.float64)
        for t in range(len(self.reward_memory)):
            G_sum = 0
            discount = 1
            for k in range(t, len(self.reward_memory)):
                G_sum += self.reward_memory[k] * discount
                discount *= self.gamma
            G[t] = G_sum
        G = T.tensor(G, dtype=T.float).to(self.policy.device)
        
        # 计算损失:-sum(G_t * log_prob(a_t|s_t))
        loss = 0
        for g, logprob in zip(G, self.log_prob_memory):
            loss += -g * logprob
        loss.backward()
        self.policy.optimizer.step()

        # 清空记忆池
        self.log_prob_memory = []
        self.reward_memory = []

# 测试Pendulum-v1环境
env = gym.make('Pendulum-v1')
n_games = 1000  
agent = PolicyGradientAgent(gamma=0.99, lr=0.001, input_dims=[3], action_dim=1)

scores = []
for i in range(n_games):
    done = False
    observation, _ = env.reset()  # 适配新版本gym的reset返回格式
    score = 0
    while not done:
        action = agent.choose_action(observation)
        observation_, reward, terminated, truncated, info = env.step([action])  # Pendulum需传入列表形式的动作
        done = terminated or truncated
        score += reward
        # env.render()  # 需要可视化可取消注释
        agent.store_rewards(reward)
        observation = observation_
    agent.learn()
    scores.append(score)
    
    # 每50轮打印一次平均结果
    if i % 50 == 0:
        avg_score = np.mean(scores[-50:])
        print(f"Episode {i}, Average Score: {avg_score:.2f}")

env.close()

细节说明

  1. 策略网络输出处理:

    • 最后一层输出2*action_dim个值,拆分后分别作为均值mu和标准差的原始输出
    • 使用softplus激活函数处理标准差,相比exp更不容易出现数值爆炸问题,同时加1e-6避免标准差为0导致的数值错误
  2. 动作采样与裁剪:

    • 基于Normal分布采样动作后,必须根据环境的动作范围进行裁剪(如Pendulum的动作范围是[-2,2]),否则环境会抛出非法动作的错误
    • 新版本gym的step方法需要传入列表形式的动作(如[action]),需注意适配
  3. 损失计算逻辑:

    • 核心损失公式和离散版本一致:loss = -Σ(G_t * log_prob(a_t|s_t)),只是对数概率来自正态分布而非分类分布

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 02:15:35