第19篇:策略梯度方法:从REINFORCE到Actor-Critic
·
摘要:
本文系统讲解策略梯度定理(Policy Gradient Theorem),深入解析REINFORCE算法及其基线(Baseline)减方差技巧,详解Actor-Critic架构(Actor网络 + Critic网络)、A2C(Advantage Actor-Critic)与A3C(异步A2C)。结合PyTorch,实战Pendulum-v1连续控制任务。帮助学习者掌握“直接优化策略”的强化学习高级方法。
一、值函数方法的局限
在DQN等值函数方法中:
- 学习
Q(s,a)或V(s)。 - 策略由值函数导出(如
π(s) = argmax_a Q(s,a))。
1.1 问题
- ❌ 离散动作空间:难以处理连续动作(如机器人关节角度)。
- ❌ 次优策略:
argmax可能忽略探索。 - ❌ 确定性策略:缺乏随机性,不利于探索。
✅ 策略梯度方法直接学习策略
π(a|s; θ),可处理连续动作与随机策略。
二、策略梯度定理(Policy Gradient Theorem)
2.1 目标函数
最大化期望累积奖励:
J(θ) = E_π [ Σ_t γ^t rₜ ]
2.2 梯度更新
策略梯度定理给出:
∇_θ J(θ) = E_π [ ∇_θ log π(a|s; θ) · Q^π(s,a) ]
∇_θ log π(a|s; θ):得分函数(Score Function),指示策略应如何调整。Q^π(s,a):优势(Advantage),动作a相对于平均表现的好坏。
✅ 梯度方向:提升好动作(高Q)的概率,降低差动作(低Q)的概率。
三、REINFORCE算法:蒙特卡洛策略梯度
3.1 核心思想
- 使用完整轨迹(Episode)的回报
Gₜ作为Q(s,a)的无偏估计。 - 更新规则:
θ ← θ + α ∇_θ log π(aₜ|sₜ; θ) Gₜ
3.2 算法步骤
- 用策略
π(θ)采样一条完整轨迹{sₜ, aₜ, rₜ}。 - 计算每个时间步的回报
Gₜ = Σ_{k=t}^T γ^{k-t} rₖ。 - 更新策略参数
θ。
3.3 问题:高方差
Gₜ是Q(s,a)的高方差估计,导致训练不稳定。
四、减方差技巧:基线(Baseline)
4.1 原理
引入基线函数 b(s),不改变梯度期望:
∇_θ J(θ) = E_π [ ∇_θ log π(a|s; θ) · (Q^π(s,a) - b(s)) ]
- 常用
b(s) = V^π(s)(状态价值函数)。 (Q - V)即为优势函数A(s,a)。
✅ 减少方差,加速收敛。
五、Actor-Critic架构:结合策略与价值
5.1 核心思想
- Actor:策略网络
π(a|s; θ),负责“行动”。 - Critic:价值网络
V(s; w),负责“评价”Actor的表现。
5.2 工作流程
- Actor根据策略选择动作
a ~ π(s; θ)。 - 环境返回
r, s'。 - Critic评估新状态,计算优势
A(s,a) ≈ r + γV(s') - V(s)。 - Actor用优势更新策略:
θ ← θ + α ∇_θ log π(a|s; θ) A(s,a)。 - Critic更新价值网络:
w ← w + β [G - V(s; w)] ∇_w V(s; w)。
✅ Actor-Critic是在线学习,无需完整轨迹。
六、A2C:Advantage Actor-Critic
6.1 优势函数估计
使用TD误差作为优势估计:
A(sₜ,aₜ) ≈ δₜ = rₜ + γV(sₜ₊₁) - V(sₜ)
δₜ:TD误差,是A(s,a)的有偏但低方差估计。
6.2 损失函数
-
Critic损失(价值网络):
L_critic = (r + γV(s') - V(s))² (MSE) -
Actor损失(策略网络):
L_actor = - log π(a|s; θ) * A(s,a)
✅ A2C是同步的,通常在单个环境中训练。
七、A3C:异步优势Actor-Critic
7.1 创新点
- 多个智能体(Worker)在多个环境副本中异步运行。
- 各Worker独立采样并计算梯度。
- 梯度汇总到全局网络进行更新。
7.2 优点
- ✅ 高效并行:加速数据采集。
- ✅ 自然正则化:不同Worker探索不同策略。
- ✅ 稳定:异步更新减少相关性。
✅ A3C在Atari游戏上取得突破性成果。
八、实战:使用A2C解决Pendulum-v1连续控制
8.1 环境介绍
- 任务:控制摆杆使其倒立。
- 状态:3维
[cos(θ), sin(θ), θ_dot] - 动作:1维连续值
[-2, 2](扭矩) - 奖励:负的角速度与角度偏差,越接近0越好。
8.2 环境准备
pip install gymnasium[box2d] torch matplotlib
8.3 Actor-Critic网络定义
import torch
import torch.nn as nn
import torch.optim as optim
import gymnasium as gym
import numpy as np
# Actor网络(高斯策略)
class Actor(nn.Module):
def __init__(self, state_dim, action_dim, max_action):
super(Actor, self).__init__()
self.fc = nn.Sequential(
nn.Linear(state_dim, 256),
nn.ReLU(),
nn.Linear(256, 256),
nn.ReLU()
)
self.mu_head = nn.Linear(256, action_dim) # 均值
self.log_std_head = nn.Linear(256, action_dim) # 对数标准差
self.max_action = max_action
def forward(self, state):
x = self.fc(state)
mu = self.max_action * torch.tanh(self.mu_head(x)) # 约束范围
log_std = self.log_std_head(x)
log_std = torch.clamp(log_std, -20, 2) # 数值稳定
return mu, log_std
def sample(self, state):
mu, log_std = self.forward(state)
std = log_std.exp()
dist = torch.distributions.Normal(mu, std)
action = dist.rsample() # 重参数化采样
log_prob = dist.log_prob(action).sum(dim=-1, keepdim=True)
# Tanh变换后的对数概率(需修正)
action = torch.tanh(action)
log_prob -= torch.log(1 - action.pow(2) + 1e-6).sum(dim=-1, keepdim=True)
return action, log_prob
# Critic网络(价值函数)
class Critic(nn.Module):
def __init__(self, state_dim):
super(Critic, self).__init__()
self.fc = nn.Sequential(
nn.Linear(state_dim, 256),
nn.ReLU(),
nn.Linear(256, 256),
nn.ReLU(),
nn.Linear(256, 1)
)
def forward(self, state):
return self.fc(state)
8.4 A2C智能体
class A2CAgent:
def __init__(self, state_dim, action_dim, max_action):
self.actor = Actor(state_dim, action_dim, max_action)
self.critic = Critic(state_dim)
self.actor_optimizer = optim.Adam(self.actor.parameters(), lr=3e-4)
self.critic_optimizer = optim.Adam(self.critic.parameters(), lr=3e-4)
self.gamma = 0.99
def select_action(self, state):
state = torch.FloatTensor(state).unsqueeze(0)
with torch.no_grad():
action, _ = self.actor.sample(state)
return action.numpy().flatten()
def update(self, state, action, reward, next_state, done):
state = torch.FloatTensor(state).unsqueeze(0)
action = torch.FloatTensor(action).unsqueeze(0)
reward = torch.FloatTensor([reward])
next_state = torch.FloatTensor(next_state).unsqueeze(0)
done = torch.BoolTensor([done])
# Critic更新
value = self.critic(state)
with torch.no_grad():
next_value = self.critic(next_state)
target_value = reward + self.gamma * next_value * (~done)
critic_loss = nn.MSELoss()(value, target_value)
self.critic_optimizer.zero_grad()
critic_loss.backward()
self.critic_optimizer.step()
# Actor更新
mu, log_std = self.actor(state)
dist = torch.distributions.Normal(mu, log_std.exp())
log_prob = dist.log_prob(action).sum(dim=-1, keepdim=True)
# 优势估计
advantage = target_value - value
actor_loss = -(log_prob * advantage.detach()).mean()
self.actor_optimizer.zero_grad()
actor_loss.backward()
self.actor_optimizer.step()
return actor_loss.item(), critic_loss.item()
8.5 训练循环
env = gym.make('Pendulum-v1')
state_dim = env.observation_space.shape[0] # 3
action_dim = env.action_space.shape[0] # 1
max_action = float(env.action_space.high[0]) # 2.0
agent = A2CAgent(state_dim, action_dim, max_action)
rewards_history = []
for episode in range(500):
state, _ = env.reset()
total_reward = 0
done = False
while not done:
action = agent.select_action(state)
next_state, reward, terminated, truncated, _ = env.step(action)
done = terminated or truncated
actor_loss, critic_loss = agent.update(state, action, reward, next_state, done)
state = next_state
total_reward += reward
rewards_history.append(total_reward)
if episode % 50 == 0:
print(f"Episode {episode}, Total Reward: {total_reward:.2f}")
# 绘制学习曲线
import matplotlib.pyplot as plt
plt.plot(rewards_history)
plt.xlabel('Episode')
plt.ylabel('Total Reward')
plt.title('A2C on Pendulum-v1')
plt.show()
✅ 智能体通常在200-300轮后将奖励稳定在-200左右(接近最优)。
九、总结与学习建议
本文我们:
- 理解了策略梯度直接优化策略的优势;
- 掌握了REINFORCE算法与基线减方差;
- 学习了Actor-Critic的“演员-评论家”协作机制;
- 深入理解了A2C与A3C;
- 实战了连续控制任务(Pendulum)。
📌 学习建议:
- 理解优势函数:它是策略更新的“信号”。
- 探索-利用:策略梯度天然支持随机策略。
- 扩展学习:PPO(近端策略优化)、SAC(软Actor-Critic)。
- 关注稳定性:Clip、熵正则化等技巧。
- 探索应用:机器人控制、自动驾驶、游戏AI。
十、下一篇文章预告
第20篇:深度强化学习进阶:从PPO到SAC
我们将深入讲解:
- 近端策略优化(PPO)的Clip机制与重要性采样
- 软Actor-Critic(SAC)的最大熵强化学习
- DDPG(深度确定性策略梯度)与TD3
- 使用Stable-Baselines3实现复杂控制任务
- 强化学习在机器人、金融、推荐系统的应用
进入深度强化学习的“高级战场”!
参考文献
- Sutton, R. S. & Barto, A. G. (2018). Reinforcement Learning: An Introduction. MIT Press.
- Mnih, V. et al. (2016). Asynchronous Methods for Deep Reinforcement Learning. ICML.
- Schulman, J. et al. (2017). Proximal Policy Optimization Algorithms. arXiv.
- Haarnoja, T. et al. (2018). Soft Actor-Critic: Off-Policy Maximum Entropy Deep Reinforcement Learning with a Stochastic Actor. ICML.
- Stable-Baselines3文档: Stable-Baselines3 Docs - Reliable Reinforcement Learning Implementations — Stable Baselines3 2.7.1a0 documentation
更多推荐
所有评论(0)