共计 2398 个字符,预计需要花费 6 分钟才能阅读完成。
背景与痛点
强化学习在复杂决策场景中(如机器人控制、游戏 AI)常面临两大核心问题:

- 样本效率低下 :传统方法(如 Q -Learning)需要大量与环境交互的样本才能学习有效策略。例如在 Atari 游戏中,DQN 可能需要数千万帧数据才能达到人类水平。
- 收敛困难 :高维状态空间和稀疏奖励会导致策略梯度方法(如 REINFORCE)出现训练震荡。我曾在一个物流路径优化项目中,观察到传统方法需要 3 周训练才收敛到次优解。
技术选型对比
| 算法 | 优势 | 局限性 | 适用场景 |
|---|---|---|---|
| DQN | 离散动作空间表现稳定 | 无法处理连续动作 | 游戏 AI、推荐系统 |
| PPO | 策略更新更平稳 | 超参数敏感 | 机器人控制 |
| AC 网络 | 天然支持连续 / 离散动作 | 需要精细调整 Actor-Critic 比例 | 金融交易、自动驾驶 |
AC 网络的核心优势在于:
- Actor 直接输出动作概率分布,避免 DQN 的 argmax 操作带来的量化误差
- Critic 提供低方差的价值估计,相比 REINFORCE 减少了策略梯度方差
核心实现细节
网络架构设计
class Actor(nn.Module):
def __init__(self, state_dim, action_dim):
super().__init__()
self.fc1 = nn.Linear(state_dim, 256)
self.fc2 = nn.Linear(256, action_dim) # 输出动作概率分布
def forward(self, state):
x = F.relu(self.fc1(state))
return torch.softmax(self.fc2(x), dim=-1)
class Critic(nn.Module):
def __init__(self, state_dim):
super().__init__()
self.fc1 = nn.Linear(state_dim, 256)
self.fc2 = nn.Linear(256, 1) # 输出状态价值
关键训练机制
- 经验回放 :
- 使用 Priority Experience Replay 缓冲池
-
采样时根据 TD 误差设置优先级
-
损失函数 :
# Actor 损失(策略梯度)advantage = returns - values.detach() policy_loss = -(log_probs * advantage).mean() # Critic 损失(价值估计)value_loss = F.mse_loss(returns, values)
完整代码示例
import torch
import torch.nn as nn
import torch.optim as optim
import numpy as np
from collections import deque
class AC_Agent:
def __init__(self, state_dim, action_dim):
self.actor = Actor(state_dim, action_dim)
self.critic = Critic(state_dim)
self.memory = deque(maxlen=10000)
def get_action(self, state):
probs = self.actor(state)
dist = torch.distributions.Categorical(probs)
return dist.sample().item()
def train(self, batch_size=64, gamma=0.99):
if len(self.memory) < batch_size:
return
# 从缓冲池采样
transitions = random.sample(self.memory, batch_size)
states, actions, rewards, next_states, dones = zip(*transitions)
# 计算 TD 目标
with torch.no_grad():
next_values = self.critic(next_states)
targets = rewards + gamma * next_values * (1 - dones)
# 更新 Critic
values = self.critic(states)
value_loss = F.mse_loss(values, targets)
# 更新 Actor
probs = self.actor(states)
log_probs = torch.log(probs.gather(1, actions))
advantage = targets - values.detach()
policy_loss = -(log_probs * advantage).mean()
# 联合优化
self.optimizer.zero_grad()
(value_loss + policy_loss).backward()
self.optimizer.step()
性能与安全性考量
- 梯度裁剪 :Critic 网络容易出现梯度爆炸
torch.nn.utils.clip_grad_norm_(self.critic.parameters(), 0.5) - 熵正则化 :防止策略过早收敛到局部最优
entropy = -torch.sum(probs * torch.log(probs), dim=-1) policy_loss -= 0.01 * entropy.mean() # 调节系数 0.01
生产环境避坑指南
- 超参数调优 :
- 学习率:Actor 通常比 Critic 小 3 -10 倍
- 折扣因子 γ:长期任务建议 0.99,短期任务 0.9
- 分布式训练 :
- 使用参数服务器架构
- 同步频率建议每 10-100 步一次
总结与思考
AC 网络特别适合需要精细控制的任务(如机械臂抓取)。在实践中可以尝试:
- 将 Critic 扩展为 Dueling 架构提升价值估计精度
- 结合 Meta-Learning 实现快速适应新环境
- 在自动驾驶中用于连续动作空间的速度控制
建议读者从 OpenAI Gym 的 Pendulum 环境开始实验,逐步扩展到更复杂场景。
正文完
