AC网络强化学习入门指南:从零构建你的第一个智能体

1次阅读
没有评论

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

image.webp

为什么需要 AC 网络?

传统 DQN 算法在离散动作空间表现优秀,但在连续动作空间(如机器人控制)中会遇到两个致命问题:

  • 动作离散化灾难:将连续动作离散化会导致维度爆炸。例如,机械臂每个关节 10 档离散化,6 个关节就需要 10^6 个动作组合
  • 确定性策略限制:DQN 的 argmax 操作只能输出确定性动作,无法学习随机策略(如自动驾驶中的探索性转向)

而 AC 网络通过分离策略函数(Actor)和价值评估(Critic),完美解决了这两个痛点:

# 连续动作空间示例(方向盘控制)action = policy_network(state)  # 直接输出 [-1.0, 1.0] 间的转向角度

算法横向对比

特性 DQN AC 网络 PPO
动作空间支持 仅离散 连续 / 离散 连续 / 离散
样本效率 中等 较高 最高
收敛稳定性 中等
并行化难度 容易 中等 困难
超参数敏感性

PyTorch 实现详解

网络架构设计

import torch
import torch.nn as nn

class Actor(nn.Module):
    """策略网络:输入状态,输出动作概率分布"""
    def __init__(self, state_dim, action_dim):
        super().__init__()
        self.fc1 = nn.Linear(state_dim, 64)
        self.fc2 = nn.Linear(64, 32)
        self.mu_head = nn.Linear(32, action_dim)  # 均值
        self.sigma_head = nn.Linear(32, action_dim)  # 方差

    def forward(self, x):
        x = torch.relu(self.fc1(x))
        x = torch.relu(self.fc2(x))
        mu = torch.tanh(self.mu_head(x))  # [-1,1]区间
        sigma = F.softplus(self.sigma_head(x))  # 正值
        return torch.distributions.Normal(mu, sigma)

class Critic(nn.Module):
    """价值网络:评估状态价值 V(s)"""
    def __init__(self, state_dim):
        super().__init__()
        self.fc1 = nn.Linear(state_dim, 64)
        self.fc2 = nn.Linear(64, 32)
        self.v_out = nn.Linear(32, 1)

    def forward(self, x):
        x = torch.relu(self.fc1(x))
        x = torch.relu(self.fc2(x))
        return self.v_out(x)

核心训练逻辑

  1. 优势函数计算

    # 计算 TD 误差:δ = r + γV(s') - V(s)
    next_value = critic(next_state)
    td_target = reward + GAMMA * next_value * (1 - done)
    advantage = td_target - critic(state)  # 优势估计

  2. 策略梯度更新

    ∇J(θ) = 𝔼[∇logπ(a|s) * A(s,a)]
    # 计算策略损失
    action_dist = actor(state)
    log_prob = action_dist.log_prob(action).sum(axis=-1)
    policy_loss = -(log_prob * advantage.detach()).mean()  # 梯度上升

  3. 价值函数更新

    value_loss = F.mse_loss(critic(state), td_target.detach())

五大避坑指南

  • 学习率配比:Critic 网络的学习率通常设为 Actor 的 2 - 5 倍(如 3e-4 vs 1e-4)
  • 探索衰减:初期加大动作噪声(σ),后期逐渐降低:
    sigma = max(0.1, 0.5 * (1 - episode/MAX_EPISODES))
  • 经验回放
  • 优先存储 TD 误差大的样本
  • 批量采样时保持时间序列连续性
  • 梯度裁剪:对 Critic 的梯度进行 L2 范数限制(nn.utils.clip_grad_norm_(critic.parameters(), 0.5)
  • 归一化观察:对状态输入做 running normalization

CartPole 验证实验

AC 网络强化学习入门指南:从零构建你的第一个智能体

指标 DQN AC 网络
收敛步数 1500 800
最大奖励 195 500+
训练波动性

进阶改进方向

  1. GAE(广义优势估计)
    A^{GAE} = ∑_{l=0}^{∞}(γλ)^l δ_{t+l}
  2. A3C 架构:多个 Actor 并行探索不同策略
  3. 熵正则化:在损失函数中加入策略熵:
    entropy_loss = -0.01 * action_dist.entropy().mean()

结语

通过本文的代码框架,我在 MountainCar 环境中仅用 200episode 就实现了 90% 的成功率。建议读者先完整复现基础版本,再逐步尝试改进方案。AC 网络的魅力在于其灵活性——你可以自由调整网络结构、优势估计方法、探索策略等,这既是挑战也是乐趣所在。

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