图解强化学习 |手算Actor-Critic

简介: Actor-Critic是一种融合策略优化(Actor)与价值评估(Critic)的强化学习算法:Actor负责选动作,Critic实时打分(如TD误差),实现单步更新、低方差、高效率,兼顾离散/连续动作空间。(239字)


Actor-Critic算法的基础认识

Actor-Critic:演员 - 评论家算法

价值算法:只估价值,不能直接优化动作策略,连续动作不好处理

纯策略梯度:更新方差大、收敛慢,要等回合结束才能更新

AC 算法:

Actor 直接优化行动策略,适配各类动作空间。

Critic 实时评估动作好坏,降低更新波动、加速收敛。

                         边交互边更新,不用等待回合完结,学习效率更高

image.gif


Actor-Critic算法的网络结构

Actor(策略网络)

目标:在状态 s 下,输出动作 a,让未来总奖励最大化。

输入:状态 s

输出:

         离散动作:动作概率分布

         连续动作:确定性动作或动作分布参数

学习方式:根据 Critic 的评价,调整策略参数,让 “高分动作” 出现的概率越来越高。

Critic(价值网络)

目标:学会给 Actor 的动作打分,判断 “这个动作在当前状态下好不好”。

输入:状态 s(或状态 + 动作 (s,a))

输出:

          状态价值 (V(s)):在状态 s 下,未来能拿到多少平均奖励;

          动作价值 (Q(s,a)):在状态 s 做动作 a,未来能拿到多少平均奖励。

学习方式:用 TD 误差(时序差分误差)更新价值估计,让打分越来越准。

image.gif


Actor-Critic算法的网络更新

Actor的更新:

目的:让 Actor 根据 Critic 的评价(TD 误差)调整动作策略,让好动作被更多选择,坏动作被更

少选择,持续优化决策。

(1)获取当前交互得到的状态 s、动作 a、奖励 r、下一状态 s'、是否结束 done。

(2)将状态 s 输入 Critic 网络,得到当前状态价值 V (s);将下一状态 s' 输入 Critic 网络,得到下

一状态价值 V (s')。

(3)Critic 计算 TD 误差:

                                     δ = r + γ * V (s') - V (s)

                         (δ 代表动作的好坏:正 = 好,负 = 差)

(4)将状态 s 输入 Actor 网络,得到策略对动作 a 的输出概率 / 得分。

(5)依据策略梯度定理更新在线 Actor 网络参数:

                       Actor 根据 Critic 算出的 TD 误差来更新:动作好就加强,动作差就减弱。

TD Target(时序差分目标)


优势函数(Advantage)


Actor 损失(策略梯度)


Critic的更新

目的:精准预估状态价值,缩小评估偏差,为 Actor 优化提供可靠评判依据(1)获取单步交互数

据:当前状态s、即时奖励r、下一状态s'、终止标记done

(2)把s、s'分别输入 Critic 网络,得到估值(V(s))与(V(s'))

(3)计算时序差分误差

(4)以误差为损失依据,反向传播更新 Critic 网络参数,不断修正估值结果

Critic 损失(价值回归)


image.gif


手算AC算法

网络的输入输出

input_dim = 6 (Acrobot 的 6 维观测空间:两个角余弦、两个角正弦、两个角速度)

output_dim = 3 (Acrobot 的 3 个离散动作:-1扭矩, 0扭矩, +1扭矩)

假设输入X是:X = [1.0, 0.0, 1.0, 0.0, 0.0, 0.0]

actor_output: [0.07, 0.93, 0.00] (给环境采样用的动作概率)

                   动作 0(向左猛踢,-1扭矩):0.07 (7%)

                   动作 1(腰部放松,0扭矩):0.93 (93%)

                   动作 2(向右猛踢,+1扭矩):0.00 (0%)

image.gif

critic_output: -98.5  状态价值

image.gif

选择动作

假设此时 Acrobot 的状态 state 如下(对应 [cos1, sin1, cos2, sin2, v1, v2])

state = [0.0, 1.0, 1.0, 0.0, 0.5, -0.2]

action_prob:tensor([[0.02, 0.85, 0.13]])

action = 1

(大概率会返回 1,但有极小可能返回 0 或 2(是一个采样过程)

所以这段代码的作用是根据当前动作,输入到Actor网络中选择出一个动作。

image.gif

模型更新

假设我们的 buffer 如下:

经验1:(s1, a1=2, r1=-1.0, ns1, d1=0)
经验2:(s2, a2=0, r2=-1.0, ns2, d2=0)
states = [
    [1,0,1,0,0,0],   # s1
    [0,1,0,1,0.1,0.1] # s2
]
actions = [2, 0]
rewards = [-1.0, -1.0]
dones = [0, 0]
gamma = 0.99 (默认,AC常用)

image.gif

网络前向输出

Actor 输出(动作概率):

action_prob = [
    [0.07, 0.93, 0.00],  # s1 对应动作概率
    [0.02, 0.85, 0.13]   # s2 对应动作概率
]

image.gif

Critic 输出(当前状态价值)

state_value = [
    [-98.5],  # V(s1)
    [-95.0]   # V(s2)
]

image.gif

Critic 输出(下一状态价值)

next_state_value = [
    [-97.0],  # V(s1')
    [-93.0]   # V(s2')
]

image.gif

手工计算


经验 1:

td_target1 = -1.0 + 0.99 * 1.0 * (-97.0)
= -1.0 - 96.03
= -97.03

image.gif

经验 2:

td_target2 = -1.0 + 0.99 * 1.0 * (-93.0)
= -1.0 - 92.07
= -93.07

image.gif

最终:

td_target = [-97.03, -93.07]

image.gif

计算 TD Delta(误差)


经验 1:

td_delta1 = -97.03 - (-98.5) = 1.47

image.gif

经验 2:

td_delta2 = -93.07 - (-95.0) = 1.93

image.gif

最终:

td_delta = [1.47, 1.93]

image.gif

正数 → 动作比预期好 → Actor 要加强这个动作

计算 log_prob(动作对数概率)

log_prob2 = log(0.02) ≈ -3.91

image.gif

log_prob = log( 选中动作的概率 )

image.gif

经验 1:动作 a=2,概率 = 0.00

log_prob1 = log(0.00) → 负无穷(这里我们换成 0.01 方便计算)
log(0.01) ≈ -4.6

image.gif

经验 2:动作 a=0,概率 = 0.02

log_prob2 = log(0.02) ≈ -3.91

image.gif

最终:

log_prob = [-4.6, -3.91]

image.gif

计算 Actor Loss


经验 1:
-4.6 * 1.47 = -6.76
经验 2:
-3.91 * 1.93 = -7.55
平均:
mean = (-6.76 -7.55) / 2 = -7.155
加负号:
actor_loss = -(-7.155) = 7.155

image.gif

计算 Critic Loss

经验 1:
(-98.5 + 97.03)² = (-1.47)² = 2.16
经验 2:
(-95.0 + 93.07)² = (-1.93)² = 3.72
平均:
critic_loss = (2.16 + 3.72) / 2 = 2.94
第六步:总 Loss
loss = actor_loss + 0.5 * critic_loss
= 7.155 + 0.5*2.94
= 7.155 + 1.47
= 8.625

image.gif

td_target      = [-97.03, -93.07]
td_delta       = [1.47, 1.93]
log_prob       = [-4.6, -3.91]
actor_loss     = 7.155
critic_loss    = 2.94
total_loss     = 8.625

image.gif

Critic 让自己的估值更准

Actor 因为 td_delta 为正,加强了刚才选择的动作

AC 只需要一步数据就能更新,不需要等回合结束

AC算法是单步更新,在线更新

AC算法核心更新仅依赖单步交互数据s,a,r,s',满足即时计算TD误差、完成参数更新的数学条件,

天生支持单步更新;传统策略梯度需整回合累计回报才能计算,无法单步更新。代码采用回合批量

更新,只是为提升训练稳定性与运算效率。

image.gif


核心点

Actor 损失为什么是:actor_loss = - (log_prob * td_delta).mean()

Actor 的任务:让 “好动作” 概率变高,让 “坏动作” 概率变低。

对数转换后,乘积求导转为加减求导,梯度计算大幅简化。最大化动作概率等价于最大化对数概率,配合负号就能转为深度学习通用的最小损失优化范式。

log_prob = torch.log( 策略网络输出的动作概率 )

image.gif

你选了这个动作,就取这个动作的概率,再取对数

例子:概率 = 0.07 → log (0.07) ≈ -2.66

作用:用来告诉网络 “我刚才选了哪个动作”。

td_delta = 实际收益 - 网络预估收益

image.gif

td_delta > 0:这个动作比想象中好!

td_delta < 0:这个动作比想象中差!

作用:评价刚才动作的好坏

为什么要把 log_prob 和 td_delta 乘起来?

loss_part = log_prob * td_delta

image.gif

本质含义:

✅ 如果动作好(td_delta 正)

log (prob) 是负数

正数 × 负数 = 负数

❌ 如果动作差(td_delta 负)

log (prob) 是负数

负数 × 负数 = 正数

为什么前面要加一个 负号?

actor_loss = - loss_part

image.gif

好动作 → 算出来是负数 → 加负号 → 变成正数 → 网络会减小它 → 让概率上升

坏动作 → 算出来是正数 → 加负号 → 变成负数 → 网络会增大它 → 让概率下降

神经网络永远做一件事:想尽办法 让 损失(loss) 变 小!

不管你怎么设计公式,网络只会:loss 越小越好 → 拼命减小 loss

好动作的损失是正数

actor_loss = - ( -0.2 ) = +0.2

image.gif

现在网络要做什么?

网络想让 loss 变小!

loss 现在是 0.2

网络想把它变成 0.1、0.05、0…

怎么才能让 loss 变小?

必须让 log_prob 变大(从 -0.1 → -0.05 → -0.01)

log_prob 变大 → 概率变大!

目录
相关文章
|
4月前
|
机器学习/深度学习 算法 机器人
图解强化学习 |手算SAC算法
SAC(Soft Actor-Critic)是最稳定、强大的连续动作强化学习算法,广泛应用于机器人控制与决策任务。其核心是最大熵强化学习:通过双Q网络抑制过估计,柔性策略网络增强探索,自适应温度系数α动态平衡利用与探索,兼顾最优性与鲁棒性。(239字)
660 1
|
4月前
|
机器学习/深度学习 人工智能 算法
图解强化学习 |手算近端策略优化算法(PPO)
PPO(近端策略优化)是当前最主流的强化学习算法,以训练稳定、上手简单、泛化性强著称。它通过Actor-Critic双网络架构,结合PPO-Clip损失函数限制策略更新幅度,并利用GAE优势估计提升样本效率,广泛应用于游戏AI、机器人控制、大模型对齐等领域。
1045 3
|
4月前
|
人工智能 运维 网络安全
我这半年实际使用过的 4 款 SSH 工具体验分享
本文精选4款主流SSH工具:轻量稳定的PuTTY、Windows全能神器MobaXterm、跨平台云同步的Termius,以及AI驱动的智能终端Aeroshell,覆盖从经典运维到AI辅助新场景,助开发者与运维人员高效管理服务器。(239字)
1658 1
我这半年实际使用过的 4 款 SSH 工具体验分享
|
4月前
|
数据采集 Python Windows
2472.一款图片批量提取工具:从文章到图库,一招搞定素材管理_创建自己的永久免费图床
公号图床图片提取工具:一键批量提取微众号文章中的所有配图,智能识别防盗链、自动去重、支持纯链接/HTML/论坛格式输出,并可实时预览、本地批量保存,直链引用,操作极简,效率跃升。
1294 3
2472.一款图片批量提取工具:从文章到图库,一招搞定素材管理_创建自己的永久免费图床
|
4月前
|
机器学习/深度学习 数据可视化 PyTorch
PyTorch深度学习实战 | 基于LSTM的时间序列预测任务
本文介绍了使用LSTM模型预测印度德里市平均温度的两个项目。项目1对温度数据进行归一化处理,采用滑动窗口法构建监督学习样本,设计5层LSTM网络结构,并详细说明了模型训练过程及评估方法。项目2在数据处理上增加了标准化和周期性特征,改进了网络架构,引入了学习率调整和早停机制优化训练过程。两个项目均通过可视化对比预测值和真实值,验证了LSTM模型在时间序列预测中的有效性。文章从数据处理、模型构建到训练优化,完整呈现了温度预测的实现流程,为时序预测任务提供了实用参考。
292 0
|
4月前
|
机器学习/深度学习 人工智能 API
5款靠谱的IP归属地查询服务深度测评:准确率、性能、离线库谁更强?
本文实测5款IP归属地查询工具,直击城市级定位不准痛点:广告投放偏差、风控失效。建议:先通过服务商提供的的免费测试额度验证区县级定位效果,用真实业务样本对比竞品差异,再决定是否接入离线库。高精度不是概念,而是可落地的工程能力。
1128 2
|
4月前
|
机器学习/深度学习 存储 人工智能
图解人工智能的数学基础(线性代数)
本文系统讲解线性代数核心概念,涵盖向量(定义、几何/坐标表示、内积)、矩阵(含义、运算、秩、逆、相似、分解)、行列式(几何意义与变换关系)、线性方程组、特征值与特征向量、二次型、向量空间及范数等,强调其在AI与神经网络中的实际应用。
533 7
|
4月前
|
机器学习/深度学习 人工智能 PyTorch
PyTorch深度学习实战 | 人工智能项目从训练到部署
本项目基于LSTM模型对污水处理厂总曝气量(旧区+新区)进行时序预测。通过数据清洗、Min-Max归一化、滑动窗口构造(12小时输入→预测未来1小时),构建并训练轻量级LSTM模型,支持API部署与实时调用,已实现端到端预测流程及模型保存。
346 6
|
4月前
|
机器学习/深度学习 算法 自动驾驶
图解强化学习 |手算DDPG
DDPG(深度确定性策略梯度)是一种面向连续动作空间的Actor-Critic强化学习算法。它采用4网络结构(Actor/Critic及其对应目标网络),结合经验回放与软更新,通过确定性策略梯度优化策略,广泛应用于机器人控制、自动驾驶等场景。(239字)
417 1
|
4月前
|
人工智能 开发工具 开发者
学习AI Agent编程-第一天-MCP基础
本文精炼解析MCP(Model Context Protocol):它不是新模型,而是让AI Agent运行时动态增删工具的协议。通过MCP Server(工具实现)、Client(SDK封装)与Host(Agent应用)三组件协作,解决传统`bind_tools`静态绑定的局限。附完整可运行示例,助你快速掌握80%核心用法。(239字)
604 1

热门文章

最新文章