共计 2866 个字符,预计需要花费 8 分钟才能阅读完成。
技术背景
强化学习(Reinforcement Learning, RL)是一种通过与环境交互来学习最优策略的机器学习方法。与监督学习不同,强化学习没有标注数据,而是通过奖励信号来指导学习过程。传统的强化学习方法如 Q -learning 和 Policy Gradient 各有优缺点:
- Q-learning:基于值函数,适合离散动作空间,但难以处理连续动作空间。
- Policy Gradient:直接优化策略,适合连续动作空间,但样本效率低且训练不稳定。
AC(Actor-Critic)网络结合了这两种方法的优点,通过双网络结构(Actor 和 Critic)实现了更高效的训练。Actor 负责选择动作,Critic 负责评估动作的价值,两者协同工作,显著提升了训练的稳定性和样本效率。
架构解析
AC 网络的核心是双网络结构:
- Actor 网络 :输出动作的概率分布,通常是一个策略函数。
- Critic 网络 :评估当前状态的价值,即状态值函数。
两者的协同工作机制如下:
- Actor 根据当前状态选择动作,并执行该动作。
- Critic 评估该动作的价值,并反馈给 Actor 以调整策略。
- 通过这种方式,Actor 和 Critic 共同优化策略,逐步逼近最优解。

代码实现
以下是一个基于 PyTorch 的 AC 网络实现,演示在 CartPole 环境中的训练过程。
import torch
import torch.nn as nn
import torch.optim as optim
import gym
import numpy as np
# 定义 Actor 网络
class Actor(nn.Module):
def __init__(self, state_dim, action_dim):
super(Actor, self).__init__()
self.fc1 = nn.Linear(state_dim, 64)
self.fc2 = nn.Linear(64, 32)
self.fc3 = nn.Linear(32, action_dim)
self.softmax = nn.Softmax(dim=-1)
def forward(self, x):
x = torch.relu(self.fc1(x))
x = torch.relu(self.fc2(x))
x = self.softmax(self.fc3(x))
return x
# 定义 Critic 网络
class Critic(nn.Module):
def __init__(self, state_dim):
super(Critic, self).__init__()
self.fc1 = nn.Linear(state_dim, 64)
self.fc2 = nn.Linear(64, 32)
self.fc3 = nn.Linear(32, 1)
def forward(self, x):
x = torch.relu(self.fc1(x))
x = torch.relu(self.fc2(x))
x = self.fc3(x)
return x
# 初始化环境和网络
env = gym.make('CartPole-v1')
state_dim = env.observation_space.shape[0]
action_dim = env.action_space.n
actor = Actor(state_dim, action_dim)
critic = Critic(state_dim)
actor_optimizer = optim.Adam(actor.parameters(), lr=0.001)
critic_optimizer = optim.Adam(critic.parameters(), lr=0.001)
# 训练循环
for episode in range(1000):
state = env.reset()
total_reward = 0
done = False
while not done:
state_tensor = torch.FloatTensor(state).unsqueeze(0)
action_probs = actor(state_tensor)
action = torch.multinomial(action_probs, 1).item()
next_state, reward, done, _ = env.step(action)
total_reward += reward
# 计算 Critic 的损失
value = critic(state_tensor)
next_value = critic(torch.FloatTensor(next_state).unsqueeze(0))
target = reward + (1 - done) * 0.99 * next_value
critic_loss = nn.MSELoss()(value, target.detach())
# 计算 Actor 的损失
advantage = target.detach() - value.detach()
action_log_prob = torch.log(action_probs.squeeze(0)[action])
actor_loss = -action_log_prob * advantage
# 更新网络
actor_optimizer.zero_grad()
critic_optimizer.zero_grad()
actor_loss.backward()
critic_loss.backward()
actor_optimizer.step()
critic_optimizer.step()
state = next_state
print(f'Episode {episode}, Total Reward: {total_reward}')
调优实践
- 超参数选择 :
- 学习率:通常设置为 0.001 左右,过大容易震荡,过小收敛慢。
- 折扣因子(gamma):控制未来奖励的权重,一般设为 0.99。
-
批量大小:影响训练的稳定性,建议从 32 开始尝试。
-
奖励函数设计 :
-
奖励函数的设计直接影响训练效果。在 CartPole 中,简单的平衡奖励即可,但在复杂环境中可能需要精心设计。
-
训练稳定性保障 :
- 使用经验回放(Replay Buffer)可以减少样本相关性,提高训练稳定性。
- 定期保存模型参数,防止训练中断导致的数据丢失。
性能分析
AC 网络在训练速度和最终表现上通常优于传统方法:
- 训练速度 :由于 Critic 提供了更准确的梯度信号,AC 网络通常比 Policy Gradient 收敛更快。
- 最终表现 :AC 网络在复杂环境中表现更稳定,能够学习到更优的策略。
避坑指南
- 梯度消失 :
- 使用 ReLU 等激活函数避免梯度消失。
-
适当调整网络深度,避免过深的网络导致梯度消失。
-
探索不足 :
- 在 Actor 的输出层添加熵正则项,鼓励探索。
- 使用 ε -greedy 策略,在训练初期增加随机动作的概率。
延伸思考
- AC 网络如何扩展到多智能体强化学习场景?
- 在连续动作空间中,AC 网络如何进一步优化?
- 如何结合 AC 网络与注意力机制以处理高维状态空间?
希望这篇博客能帮助你快速掌握 AC 网络强化学习的核心原理和工程实践。如果有任何问题或建议,欢迎在评论区留言讨论!
正文完
