跳到主要内容
极客日志极客日志面向AI+效率的开发者社区
首页博客GitHub 精选镜像AI 生图工具UI配色美学隐私政策关于联系
搜索内容 / 工具 / 仓库 / 镜像...⌘K搜索
注册
博客列表
PythonAI算法

强化学习:演员评论家 Actor-Critic 算法详解与实现

Actor-Critic 算法结合策略梯度与值函数估计,通过 Actor 优化动作选择,Critic 评估状态价值以降低方差。本文详解其数学原理、算法流程及 PyTorch 实战代码,涵盖 CartPole 环境训练示例,帮助理解强化学习中策略优化的核心机制。

萤火微光发布于 2026/3/27更新于 2026/7/2330 浏览
强化学习:演员评论家 Actor-Critic 算法详解与实现

演员评论家 Actor-Critic 算法

Actor-Critic 算法是强化学习中一种经典且高效的方法,它巧妙地将策略梯度(Policy Gradient)和值函数估计(Value Function Estimation)结合起来。简单来说,它通过两个网络协同工作:一个负责行动(Actor),另一个负责评价(Critic)。

核心概念理解

1. 角色设定

想象你正在训练一个机器人爬山,目标是到达山顶获得最高奖励。

  • Actor(演员):相当于冒险家,负责根据当前状态决定下一步怎么走(选择动作)。它的策略可能不完美,需要不断调整。
  • Critic(评论家):相当于导师,站在旁边观察并评估 Actor 的表现。它会告诉 Actor:'这一步走得好'或者'离目标更远了'。

2. 协作机制

两者的协作过程非常直观:

  1. Actor 观察环境状态,根据当前策略选择一个动作。
  2. 环境反馈新的状态和奖励。
  3. Critic 根据反馈计算价值(Value),评估刚才那个动作的好坏。
  4. Actor 利用 Critic 的评价来更新自己的策略参数,让下一次选择更优。

这种分工使得 Actor 专注于决策,而 Critic 专注于提供低方差的梯度信号,两者互补,比单独使用任何一种方法效果更好。

3. 为什么叫 Actor-Critic?

名字直接反映了功能分工:Actor 负责执行动作,Critic 负责评判价值。结合后的优势在于,Critic 提供的基准线(Baseline)可以有效降低策略梯度的方差,从而加速收敛。

算法背景与数学推导

1. 优化目标

强化学习的核心目标是最大化累积折扣奖励的期望: $$J(\theta) = \mathbb{E}{\pi\theta} \left[ \sum_{t=0}^\infty \gamma^t r_t \right]$$ 其中 $\gamma$ 是折扣因子,$r_t$ 是即时奖励,$\pi_\theta(a|s)$ 是策略函数。

2. 策略梯度定理

为了优化策略参数 $\theta$,我们计算目标函数的梯度: $$\nabla_\theta J(\theta) = \mathbb{E}{\pi\theta} \left[ \nabla_\theta \log \pi_\theta(a|s) \cdot A^\pi(s, a) \right]$$ 这里的 $A^\pi(s, a)$ 是优势函数,衡量动作相对于平均水平的优劣。直接使用环境反馈会导致高方差问题,这就是引入 Critic 的原因。

3. Critic:值函数估计

Critic 的目标是学习状态值函数 $V^\pi(s)$,通常用神经网络近似。它通过最小化时间差分(TD)误差来更新: $$\delta_t = r_t + \gamma V^\pi(s_{t+1}) - V^\pi(s_t)$$ Critic 的网络参数 $w$ 更新方向为: $$\nabla_w L(w) = \delta_t \cdot \nabla_w V^\pi(s)$$

4. Actor:策略优化

Actor 利用 Critic 计算的 TD 误差 $\delta$ 来指导策略更新: $$\theta \leftarrow \theta + \alpha \cdot \nabla_\theta \log \pi_\theta(a|s) \cdot \delta$$ 这意味着如果 Critic 认为某个动作比预期好($\delta > 0$),Actor 就增加该动作的概率;反之则减少。

实战代码实现

下面是一个基于 PyTorch 的完整 Actor-Critic 实现示例。我们将分别构建策略网络(Actor)和价值网络(Critic),并在 CartPole 环境中进行训练。

网络结构定义

首先定义 Actor 和 Critic 的网络类。Actor 输出动作概率分布,Critic 输出状态价值。

import torch
from torch import nn
from torch.nn  functional  F
 numpy  np

 (nn.Module):
    
     ():
        (PolicyNet, ).__init__()
        .fc1 = nn.Linear(n_states, n_hiddens)
        .fc2 = nn.Linear(n_hiddens, n_actions)

     ():
        x = .fc1(x)
        x = F.relu(x)
        x = .fc2(x)
        
         F.softmax(x, dim=)

 (nn.Module):
    
     ():
        (ValueNet, ).__init__()
        .fc1 = nn.Linear(n_states, n_hiddens)
        .fc2 = nn.Linear(n_hiddens, )

     ():
        x = .fc1(x)
        x = F.relu(x)
        x = .fc2(x)
         x
import
as
import
as
class
PolicyNet
"""Actor: 策略网络"""
def
__init__
self, n_states, n_hiddens, n_actions
super
self
self
self
def
forward
self, x
self
self
# 输出每个动作的概率
return
1
class
ValueNet
"""Critic: 价值网络"""
def
__init__
self, n_states, n_hiddens
super
self
self
self
1
def
forward
self, x
self
self
return

算法主逻辑

这里实现了核心的 update 方法。注意我们在计算 Loss 时,Actor 的 Loss 依赖于 Critic 输出的 TD 误差,而 Critic 的 Loss 则是预测值与目标的均方误差。

class ActorCritic:
    def __init__(self, n_states, n_hiddens, n_actions, actor_lr, critic_lr, gamma):
        self.gamma = gamma
        self.actor = PolicyNet(n_states, n_hiddens, n_actions)
        self.critic = ValueNet(n_states, n_hiddens)
        
        self.actor_optimizer = torch.optim.Adam(self.actor.parameters(), lr=actor_lr)
        self.critic_optimizer = torch.optim.Adam(self.critic.parameters(), lr=critic_lr)

    def take_action(self, state):
        """根据当前策略采样动作"""
        state = torch.tensor(state[np.newaxis, :], dtype=torch.float)
        probs = self.actor(state)
        action_dist = torch.distributions.Categorical(probs)
        action = action_dist.sample().item()
        return action

    def update(self, transition_dict):
        """批量更新模型参数"""
        states = torch.tensor(transition_dict['states'], dtype=torch.float)
        actions = torch.tensor(transition_dict['actions']).view(-1, 1)
        rewards = torch.tensor(transition_dict['rewards'], dtype=torch.float).view(-1, 1)
        next_states = torch.tensor(transition_dict['next_states'], dtype=torch.float)
        dones = torch.tensor(transition_dict['dones'], dtype=torch.float).view(-1, 1)

        # Critic 计算 TD 误差
        td_value = self.critic(states)
        td_target = rewards + self.gamma * self.critic(next_states) * (1 - dones)
        td_delta = td_target - td_value

        # Actor 损失:负对数似然 * TD 误差
        log_probs = torch.log(self.actor(states).gather(1, actions))
        actor_loss = torch.mean(-log_probs * td_delta.detach())

        # Critic 损失:均方误差
        critic_loss = torch.mean(F.mse_loss(self.critic(states), td_target.detach()))

        # 反向传播与更新
        self.actor_optimizer.zero_grad()
        self.critic_optimizer.zero_grad()
        actor_loss.backward()
        critic_loss.backward()
        self.actor_optimizer.step()
        self.critic_optimizer.step()

训练流程

最后,我们在 CartPole 环境中运行训练循环。这里需要注意数据收集的方式,我们通常在一个 episode 结束后再统一更新一次参数,这样能利用整个轨迹的信息。

import gym
import matplotlib.pyplot as plt

env_name = 'CartPole-v1'
num_episodes = 100
gamma = 0.9
actor_lr = 1e-3
critic_lr = 1e-2
n_hiddens = 16

env = gym.make(env_name)
n_states = env.observation_space.shape[0]
n_actions = env.action_space.n

agent = ActorCritic(
    n_states=n_states,
    n_hiddens=n_hiddens,
    n_actions=n_actions,
    actor_lr=actor_lr,
    critic_lr=critic_lr,
    gamma=gamma
)

return_list = []

for i in range(num_episodes):
    state = env.reset()[0]
    done = False
    episode_return = 0
    transition_dict = {
        'states': [], 'actions': [], 'next_states': [],
        'rewards': [], 'dones': []
    }

    while not done:
        action = agent.take_action(state)
        next_state, reward, done, _, _ = env.step(action)
        
        transition_dict['states'].append(state)
        transition_dict['actions'].append(action)
        transition_dict['next_states'].append(next_state)
        transition_dict['rewards'].append(reward)
        transition_dict['dones'].append(done)
        
        state = next_state
        episode_return += reward

    return_list.append(episode_return)
    agent.update(transition_dict)
    print(f'iter:{i}, return:{np.mean(return_list[-10:])}')

plt.plot(return_list)
plt.title('Training Return')
plt.show()

关键点总结

在实际应用中,有几个细节值得注意:

  1. Critic 的稳定性:Critic 的误差直接决定了 Actor 的梯度方向,如果 Critic 学得太慢或太偏,Actor 可能会学歪。
  2. 熵正则化:为了防止策略过早收敛到局部最优,可以在 Actor 的损失中加入熵项鼓励探索。
  3. 超参数敏感:Actor 和 Critic 的学习率通常需要分开调节,有时 Critic 需要更快的学习速度。
  4. 扩展性:基础的 Actor-Critic 可以进一步扩展为 A3C(异步多进程)或 PPO(近端策略优化),以解决更复杂的任务。

总结

Actor-Critic 算法通过结合策略梯度的灵活性和值函数的低方差特性,成为强化学习领域的重要基石。从数学原理到 PyTorch 落地,理解其核心在于把握 Actor 与 Critic 之间的梯度传递关系。希望这篇教程能帮助你更好地掌握这一算法,并在实际项目中灵活运用。

目录

  1. 演员评论家 Actor-Critic 算法
  2. 核心概念理解
  3. 1. 角色设定
  4. 2. 协作机制
  5. 3. 为什么叫 Actor-Critic?
  6. 算法背景与数学推导
  7. 1. 优化目标
  8. 2. 策略梯度定理
  9. 3. Critic:值函数估计
  10. 4. Actor:策略优化
  11. 实战代码实现
  12. 网络结构定义
  13. 算法主逻辑
  14. 训练流程
  15. 关键点总结
  16. 总结
  • 免费图片AI生成工具免费生成了解详情
  • Magick API 一键接入全球大模型注册送1000万token查看
  • 免费图片视频在线生成30秒,将你的创意变成现实开始设计
  • X/Twitter免费视频下载器免登陆无限额度免费视频解析下载了解详情
  • 100+免费在线小游戏爽一把
极客日志微信公众号二维码

微信扫一扫,关注极客日志

微信公众号「极客日志V2」,在微信中扫描左侧二维码关注。展示文案:极客日志V2 zeeklog

更多推荐文章

查看全部
  • C++ 二叉搜索树简单实现:增删查改全攻略
  • JNI 开发:C++ Debug 正常 Release 返回 NaN 原因解析
  • GitHub Copilot Pro 学生免费权益获取与 VS Code 配置指南
  • OpenCode 开源 AI 编程助手:从入门到精通
  • Logseq 本地部署与 cpolar 远程访问配置指南
  • 基于 YOLO 与 LLM 的 Web 目标检测与智能分析系统
  • Web 团队开发 App:是否应该选择 Capacitor
  • Whisper.cpp 离线语音识别快速入门
  • Linux 进程控制详解
  • OpenClaw 龙虾机器人 Windows 系统部署指南
  • BR8654 蓝牙 6.0 SOC 芯片技术规格与特性
  • 文心一言 4.0 调用性能优化实战
  • ClawX 可视化 AI 智能体使用指南
  • Claude Code 跨平台安装与配置指南(Win/Linux/Mac)
  • OpenClaw Web Search 配置与渠道选择指南
  • C++ 嵌入 Python 调用实战:Py_Initialize 初始化与函数交互
  • 大模型核心组件解析:激活函数与 FFN 块详解
  • GESP 2025 年 12 月 C++ 五级认证真题解析(单选 1-15)
  • 龙年 AI 生成封面图片玩法与变现指南
  • AMD 显卡部署 ComfyUI-Zluda 实战指南

相关免费在线工具

  • 加密/解密文本

    使用加密算法(如AES、TripleDES、Rabbit或RC4)加密和解密文本明文。 在线工具,加密/解密文本在线工具,online

  • RSA密钥对生成器

    生成新的随机RSA私钥和公钥pem证书。 在线工具,RSA密钥对生成器在线工具,online

  • Mermaid 预览与可视化编辑

    基于 Mermaid.js 实时预览流程图、时序图等图表,支持源码编辑与即时渲染。 在线工具,Mermaid 预览与可视化编辑在线工具,online

  • 随机西班牙地址生成器

    随机生成西班牙地址(支持马德里、加泰罗尼亚、安达卢西亚、瓦伦西亚筛选),支持数量快捷选择、显示全部与下载。 在线工具,随机西班牙地址生成器在线工具,online

  • Gemini 图片去水印

    基于开源反向 Alpha 混合算法去除 Gemini/Nano Banana 图片水印,支持批量处理与下载。 在线工具,Gemini 图片去水印在线工具,online

  • curl 转代码

    解析常见 curl 参数并生成 fetch、axios、PHP curl 或 Python requests 示例代码。 在线工具,curl 转代码在线工具,online