AC强化学习入门指南:从零构建智能决策系统

1次阅读
没有评论

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

image.webp

强化学习算法对比

在开始 AC 强化学习之前,我们先看看它与其他经典算法的区别:

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)

性能优化技巧

探索策略对比

  1. ε-greedy
  2. 简单易实现
  3. 适合离散动作空间
  4. 探索效率较低

  5. OU 噪声

  6. 适合连续动作空间
  7. 具有时间相关性
  8. 参数调优复杂

推荐初始阶段使用 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 过估计问题

症状:价值评估持续偏高导致策略失效
解决:

  1. 使用目标网络延迟更新
  2. 双 Q 学习技术
  3. 限制 Critic 输出范围

策略震荡识别

早期特征:

  • 策略熵值剧烈波动
  • 平均奖励无法收敛
  • 梯度范数异常增大

处理方法:

  1. 增加批次大小
  2. 降低学习率
  3. 添加策略熵正则项

开放性问题

  1. 对于机械臂控制等高维动作空间,如何改进 AC 框架?
  2. 分层策略设计
  3. 动作空间分解
  4. 注意力机制引入

  5. 分布式训练架构需要考虑:

  6. 参数服务器设计
  7. 梯度聚合频率
  8. 异步更新策略

希望这篇指南能帮助你快速入门 AC 强化学习。在实际项目中,建议先从简单环境(如 CartPole)开始验证算法正确性,再逐步迁移到复杂场景。

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