共计 2422 个字符,预计需要花费 7 分钟才能阅读完成。
为什么选择 A2C?
在深度强化学习(Deep Reinforcement Learning, DRL)领域,A2C(Advantage Actor-Critic)算法就像是一个平衡型的选手。相比 DQN(Deep Q-Network)这种纯价值学习的方法,A2C 同时学习策略(Actor)和价值函数(Critic),既能避免 DQN 因最大化偏差(maximization bias)导致的过估计问题,又比 PPO(Proximal Policy Optimization)这类复杂算法更易于实现。

- 与 DQN 对比:A2C 直接输出动作概率分布,适合高维 / 连续动作空间;DQN 需要维护 Q 表,离散动作空间更有优势
- 与 PPO 对比:A2C 使用简单的策略梯度,PPO 通过剪切(clip)机制约束策略更新幅度,训练更稳定但实现复杂
算法核心实现
1. 网络结构设计
A2C 的神经网络通常采用共享底层 + 独立输出头的结构:
class A2CNetwork(nn.Module):
def __init__(self, obs_dim, act_dim):
super().__init__()
# 共享特征提取层
self.shared_layers = nn.Sequential(nn.Linear(obs_dim, 64),
nn.ReLU(),
nn.Linear(64, 64),
nn.ReLU())
# 策略头(Actor)self.actor_head = nn.Linear(64, act_dim)
# 价值头(Critic)self.critic_head = nn.Linear(64, 1)
2. Advantage 计算推导
Advantage 函数的核心思想是评估当前动作的相对优势:
$$
A(s_t,a_t) = Q(s_t,a_t) – V(s_t)
$$
实际实现时常用 TD 误差(Temporal Difference error)近似:
$$
A(s_t,a_t) ≈ r_t + γV(s_{t+1}) – V(s_t)
$$
3. 并行环境采样
通过 Python 的 multiprocessing 模块实现数据并行采集:
from multiprocessing import Process, Queue
def worker(env_func, queue, steps):
env = env_func()
obs = env.reset()
for _ in range(steps):
act = model.select_action(obs)
next_obs, rew, done, _ = env.step(act)
queue.put((obs, act, rew, done, next_obs))
obs = next_obs if not done else env.reset()
完整训练代码
import torch.optim as optim
# 超参数设置
gamma = 0.99 # 折扣因子
entropy_coef = 0.01 # 熵正则项权重
lr = 7e-4 # 学习率
max_grad_norm = 0.5 # 梯度裁剪阈值
optimizer = optim.Adam(model.parameters(), lr=lr)
for epoch in range(1000):
# 1. 并行采集数据
samples = collect_samples(envs, model, 5) # 每个环境跑 5 步
# 2. 计算 Advantage 和回报
returns = compute_returns(samples['rewards'], gamma)
values = model.critic(samples['obs'])
advantages = returns - values.detach()
# 3. 策略梯度更新
log_probs = get_log_prob(model.actor(samples['obs']), samples['acts'])
actor_loss = -(log_probs * advantages).mean()
# 4. 价值函数更新
critic_loss = F.mse_loss(values, returns)
# 5. 熵正则项
entropy = get_entropy(model.actor(samples['obs']))
# 综合损失
total_loss = actor_loss + 0.5*critic_loss - entropy_coef*entropy
# 梯度裁剪
torch.nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)
# 记录到 tensorboard
writer.add_scalar('Loss/actor', actor_loss, epoch)
writer.add_scalar('Loss/critic', critic_loss, epoch)
避坑指南
- 超参数组合:
- 学习率建议从 3e- 4 到 7e- 4 尝试
-
折扣因子 γ 通常取 0.9-0.99,环境随机性越大 γ 应越小
-
梯度爆炸预防:
- 梯度裁剪阈值建议 0.5-1.0
-
可添加梯度范数监控:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.5) -
稀疏奖励应对:
- 设计合理的 reward shaping(如 CartPole 中给存活每步 +0.1)
- 使用 n -step return(建议 n =5-10)代替单步 TD
拓展思考
- 如何升级到 A3C:
- 改为异步更新(各 worker 独立计算梯度后异步更新全局模型)
-
增加模型同步频率(如每 10 次迭代同步一次)
-
连续动作空间改造:
- 策略头输出高斯分布的均值和方差
- 采用重参数化(reparameterization)技巧采样动作
- 参考 SAC(Soft Actor-Critic)的自动熵调节机制
结语
实现 A2C 就像学习骑自行车——开始可能会因为梯度不稳定而『摔跤』,但通过调整超参数、添加适当的正则化,你会发现这个算法在许多任务上都能取得不错的效果。建议先从 CartPole 这类简单环境开始验证,再逐步挑战更复杂的场景。
