共计 3741 个字符,预计需要花费 10 分钟才能阅读完成。
强化学习算法对比
在开始 AC 强化学习之前,我们先看看它与其他经典算法的区别:

| 算法类型 | 代表算法 | 更新方式 | 输出内容 | 适用场景 |
|---|---|---|---|---|
| 值函数方法 | Q-Learning | 通过 TD 误差更新 Q 表 | 状态 - 动作价值 | 离散动作空间 |
| 策略梯度方法 | Policy Gradient | 直接优化策略函数梯度 | 动作概率分布 | 连续动作空间 |
| 混合方法 | Actor-Critic | 策略梯度 + 值函数评估 | 策略 + 状态价值评估 | 连续 / 离散空间通用 |
核心原理拆解
1. Actor 网络的工作原理
Actor 网络负责生成策略 π(a|s),其输出是动作空间的概率分布:
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 = F.relu(self.fc1(x))
x = F.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)
2. Critic 网络的价值评估
Critic 网络评估状态价值 V(s),采用 MSE 损失函数:
class Critic(nn.Module):
def __init__(self, state_dim):
super().__init__()
self.fc1 = nn.Linear(state_dim, 64)
self.fc2 = nn.Linear(64, 32)
self.value_head = nn.Linear(32, 1)
def forward(self, x):
x = F.relu(self.fc1(x))
x = F.relu(self.fc2(x))
return self.value_head(x)
3. 优势函数计算
优势函数 A(s,a) = Q(s,a) – V(s),实际实现时常用 TD 误差替代:
def compute_advantage(rewards, values, next_values, dones, gamma=0.99):
# rewards: 当前步奖励
# values: 当前状态价值
# next_values: 下一步状态价值
# dones: 是否终止标志
td_errors = rewards + gamma * next_values * (1 - dones) - values
return td_errors.detach() # 阻断梯度传播
完整 PyTorch 实现
import torch
import torch.nn as nn
import torch.optim as optim
import numpy as np
from collections import deque
import random
class AC_Agent:
def __init__(self, state_dim, action_dim):
# 初始化网络
self.actor = Actor(state_dim, action_dim)
self.critic = Critic(state_dim)
# 目标网络(延迟更新)
self.target_actor = Actor(state_dim, action_dim)
self.target_critic = Critic(state_dim)
self.target_actor.load_state_dict(self.actor.state_dict())
self.target_critic.load_state_dict(self.critic.state_dict())
# 经验回放缓冲区
self.buffer = deque(maxlen=10000)
# 优化器设置梯度裁剪
self.actor_optim = optim.Adam(self.actor.parameters(), lr=1e-4)
self.critic_optim = optim.Adam(self.critic.parameters(), lr=1e-3)
def store_transition(self, state, action, reward, next_state, done):
self.buffer.append((state, action, reward, next_state, done))
def update(self, batch_size=64):
if len(self.buffer) < batch_size:
return
# 随机采样批次数据
batch = random.sample(self.buffer, batch_size)
states, actions, rewards, next_states, dones = zip(*batch)
# 转换为 Tensor
states = torch.FloatTensor(np.array(states))
actions = torch.FloatTensor(np.array(actions))
rewards = torch.FloatTensor(np.array(rewards))
next_states = torch.FloatTensor(np.array(next_states))
dones = torch.FloatTensor(np.array(dones))
# Critic 更新
current_values = self.critic(states)
next_values = self.target_critic(next_states)
advantages = compute_advantage(rewards, current_values, next_values, dones)
critic_loss = advantages.pow(2).mean()
self.critic_optim.zero_grad()
critic_loss.backward()
# 梯度裁剪防止爆炸
nn.utils.clip_grad_norm_(self.critic.parameters(), 1.0)
self.critic_optim.step()
# Actor 更新
dist = self.actor(states)
log_probs = dist.log_prob(actions)
actor_loss = -(log_probs * advantages.detach()).mean()
self.actor_optim.zero_grad()
actor_loss.backward()
nn.utils.clip_grad_norm_(self.actor.parameters(), 1.0)
self.actor_optim.step()
# 软更新目标网络
self.soft_update(self.target_actor, self.actor, tau=0.01)
self.soft_update(self.target_critic, self.critic, tau=0.01)
def soft_update(self, target, source, tau):
for target_param, param in zip(target.parameters(), source.parameters()):
target_param.data.copy_(tau*param.data + (1.0-tau)*target_param.data)
性能优化技巧
探索策略对比
- ε-greedy:
- 简单易实现
- 适合离散动作空间
-
探索效率较低
-
OU 噪声:
- 适合连续动作空间
- 具有时间相关性
- 参数调优复杂
推荐初始阶段使用 OU 噪声,后期逐渐减小噪声强度。
学习率衰减
# 在训练循环中添加
actor_scheduler = optim.lr_scheduler.StepLR(self.actor_optim, step_size=1000, gamma=0.9)
critic_scheduler = optim.lr_scheduler.StepLR(self.critic_optim, step_size=1000, gamma=0.9)
常见问题解决方案
Critic 过估计问题
症状:价值评估持续偏高导致策略失效
解决:
- 使用目标网络延迟更新
- 双 Q 学习技术
- 限制 Critic 输出范围
策略震荡识别
早期特征:
- 策略熵值剧烈波动
- 平均奖励无法收敛
- 梯度范数异常增大
处理方法:
- 增加批次大小
- 降低学习率
- 添加策略熵正则项
开放性问题
- 对于机械臂控制等高维动作空间,如何改进 AC 框架?
- 分层策略设计
- 动作空间分解
-
注意力机制引入
-
分布式训练架构需要考虑:
- 参数服务器设计
- 梯度聚合频率
- 异步更新策略
希望这篇指南能帮助你快速入门 AC 强化学习。在实际项目中,建议先从简单环境(如 CartPole)开始验证算法正确性,再逐步迁移到复杂场景。
正文完
发表至: 人工智能
近一天内
