深度强化学习实战:从零实现A2C算法及其在游戏AI中的应用

1次阅读
没有评论

共计 2422 个字符,预计需要花费 7 分钟才能阅读完成。

image.webp

为什么选择 A2C?

在深度强化学习(Deep Reinforcement Learning, DRL)领域,A2C(Advantage Actor-Critic)算法就像是一个平衡型的选手。相比 DQN(Deep Q-Network)这种纯价值学习的方法,A2C 同时学习策略(Actor)和价值函数(Critic),既能避免 DQN 因最大化偏差(maximization bias)导致的过估计问题,又比 PPO(Proximal Policy Optimization)这类复杂算法更易于实现。

深度强化学习实战:从零实现 A2C 算法及其在游戏 AI 中的应用

  • 与 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)

避坑指南

  1. 超参数组合
  2. 学习率建议从 3e- 4 到 7e- 4 尝试
  3. 折扣因子 γ 通常取 0.9-0.99,环境随机性越大 γ 应越小

  4. 梯度爆炸预防

  5. 梯度裁剪阈值建议 0.5-1.0
  6. 可添加梯度范数监控:torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.5)

  7. 稀疏奖励应对

  8. 设计合理的 reward shaping(如 CartPole 中给存活每步 +0.1)
  9. 使用 n -step return(建议 n =5-10)代替单步 TD

拓展思考

  1. 如何升级到 A3C
  2. 改为异步更新(各 worker 独立计算梯度后异步更新全局模型)
  3. 增加模型同步频率(如每 10 次迭代同步一次)

  4. 连续动作空间改造

  5. 策略头输出高斯分布的均值和方差
  6. 采用重参数化(reparameterization)技巧采样动作
  7. 参考 SAC(Soft Actor-Critic)的自动熵调节机制

结语

实现 A2C 就像学习骑自行车——开始可能会因为梯度不稳定而『摔跤』,但通过调整超参数、添加适当的正则化,你会发现这个算法在许多任务上都能取得不错的效果。建议先从 CartPole 这类简单环境开始验证,再逐步挑战更复杂的场景。

正文完
 0
评论(没有评论)