Actor-Critic 算法理论推导

1. 符号定义

在开始推导之前,我们先明确以下符号:

符号 含义
π(a∣s;θ)\pi(a|s; \boldsymbol{\theta})π(as;θ) 策略网络(Actor):给定状态 sss,输出动作 aaa 的概率分布,参数为 θ\boldsymbol{\theta}θ
q(s,a;w)q(s,a; \mathbf{w})q(s,a;w) 价值网络(Critic):评估状态-动作对 (s,a)(s,a)(s,a) 的价值,参数为 w\mathbf{w}w
V(s;θ,w)V(s; \boldsymbol{\theta}, \mathbf{w})V(s;θ,w) 状态价值函数V(s;θ,w)=∑aπ(a∣s;θ)⋅q(s,a;w)V(s; \boldsymbol{\theta}, \mathbf{w}) = \sum_a \pi(a|s; \boldsymbol{\theta}) \cdot q(s,a; \mathbf{w})V(s;θ,w)=aπ(as;θ)q(s,a;w)
dπd^{\pi}dπ 策略 π\piπ 下的状态分布
αθ,αw\alpha_{\theta}, \alpha_wαθ,αw Actor 和 Critic 的学习率

2. Actor 的更新推导(策略梯度)

2.1 优化目标

Actor 的目标是最大化期望累积回报,即最大化状态价值函数的期望:

J(θ)=Es∼dπ[V(s;θ,w)] J(\boldsymbol{\theta}) = \mathbb{E}_{s \sim d^{\pi}} [V(s; \boldsymbol{\theta}, \mathbf{w})] J(θ)=Esdπ[V(s;θ,w)]

其中 dπd^{\pi}dπ 是策略 π\piπ 下的状态分布。

2.2 策略梯度推导

我们的目标是计算 ∂V(s;θ)∂θ\frac{\partial V(s; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}}θV(s;θ),以便进行梯度上升。

步骤 1:展开状态价值函数

根据定义:
V(s;θ,w)=∑aπ(a∣s;θ)⋅q(s,a;w) V(s; \boldsymbol{\theta}, \mathbf{w}) = \sum_a \pi(a|s; \boldsymbol{\theta}) \cdot q(s,a; \mathbf{w}) V(s;θ,w)=aπ(as;θ)q(s,a;w)

θ\boldsymbol{\theta}θ 求导:
∂V(s;θ)∂θ=∂∂θ[∑aπ(a∣s;θ)⋅q(s,a;w)] \frac{\partial V(s; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}} = \frac{\partial}{\partial \boldsymbol{\theta}} \left[ \sum_a \pi(a|s; \boldsymbol{\theta}) \cdot q(s,a; \mathbf{w}) \right] θV(s;θ)=θ[aπ(as;θ)q(s,a;w)]

步骤 2:应用乘积法则

∂V(s;θ)∂θ=∑a[∂π(a∣s;θ)∂θ⋅q(s,a;w)+π(a∣s;θ)⋅∂q(s,a;w)∂θ] \begin{aligned} \frac{\partial V(s; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}} &= \sum_a \left[ \frac{\partial \pi(a|s; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}} \cdot q(s,a; \mathbf{w}) + \pi(a|s; \boldsymbol{\theta}) \cdot \frac{\partial q(s,a; \mathbf{w})}{\partial \boldsymbol{\theta}} \right] \end{aligned} θV(s;θ)=a[θπ(as;θ)q(s,a;w)+π(as;θ)θq(s,a;w)]

步骤 3:简化表达式

由于 Critic 网络 q(s,a;w)q(s,a; \mathbf{w})q(s,a;w) 的参数是 w\mathbf{w}w不依赖于 θ\boldsymbol{\theta}θ,因此:
∂q(s,a;w)∂θ=0 \frac{\partial q(s,a; \mathbf{w})}{\partial \boldsymbol{\theta}} = 0 θq(s,a;w)=0

所以:
∂V(s;θ)∂θ=∑a∂π(a∣s;θ)∂θ⋅q(s,a;w) \frac{\partial V(s; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}} = \sum_a \frac{\partial \pi(a|s; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}} \cdot q(s,a; \mathbf{w}) θV(s;θ)=aθπ(as;θ)q(s,a;w)

步骤 4:应用对数导数技巧(Log-Derivative Trick)

利用恒等式:
∂π(a∣s;θ)∂θ=π(a∣s;θ)⋅∂log⁡π(a∣s;θ)∂θ \frac{\partial \pi(a|s; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}} = \pi(a|s; \boldsymbol{\theta}) \cdot \frac{\partial \log \pi(a|s; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}} θπ(as;θ)=π(as;θ)θlogπ(as;θ)

代入得:
∂V(s;θ)∂θ=∑aπ(a∣s;θ)⋅∂log⁡π(a∣s;θ)∂θ⋅q(s,a;w) \frac{\partial V(s; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}} = \sum_a \pi(a|s; \boldsymbol{\theta}) \cdot \frac{\partial \log \pi(a|s; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}} \cdot q(s,a; \mathbf{w}) θV(s;θ)=aπ(as;θ)θlogπ(as;θ)q(s,a;w)

步骤 5:转换为期望形式

注意到上式是对所有动作 aaa 的加权求和,权重为 π(a∣s;θ)\pi(a|s; \boldsymbol{\theta})π(as;θ),这正是期望的定义:
∂V(s;θ)∂θ=EA∼π(⋅∣s;θ)[∂log⁡π(A∣s;θ)∂θ⋅q(s,A;w)] \frac{\partial V(s; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}} = \mathbb{E}_{A \sim \pi(\cdot|s; \boldsymbol{\theta})} \left[ \frac{\partial \log \pi(A|s; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}} \cdot q(s,A; \mathbf{w}) \right] θV(s;θ)=EAπ(s;θ)[θlogπ(As;θ)q(s,A;w)]

2.3 Actor 更新规则

在实际训练中,我们使用采样的动作 ata_tat 来近似期望,得到 Actor 的参数更新规则:

θ←θ+αθ⋅∂log⁡π(at∣st;θ)∂θ⋅q(st,at;w) \boxed{\boldsymbol{\theta} \leftarrow \boldsymbol{\theta} + \alpha_{\theta} \cdot \frac{\partial \log \pi(a_t|s_t; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}} \cdot q(s_t,a_t; \mathbf{w})} θθ+αθθlogπ(atst;θ)q(st,at;w)

直观理解

  • q(st,at;w)q(s_t,a_t; \mathbf{w})q(st,at;w) 越大,说明动作 ata_tat 越好,我们增大选择该动作的概率
  • q(st,at;w)q(s_t,a_t; \mathbf{w})q(st,at;w) 越小,说明动作 ata_tat 越差,我们减小选择该动作的概率

3. Critic 的更新推导(TD Learning)

3.1 时序差分学习框架

Critic 的任务是评估状态-动作对的价值。我们使用时序差分(Temporal Difference, TD)方法来更新价值估计。

定义:

  • 当前估计值qt=q(st,at;w)q_t = q(s_t, a_t; \mathbf{w})qt=q(st,at;w)
  • TD 目标yt=rt+γ⋅q(st+1,at+1;w)y_t = r_t + \gamma \cdot q(s_{t+1}, a_{t+1}; \mathbf{w})yt=rt+γq(st+1,at+1;w)
  • TD 误差δt=yt−qt=rt+γ⋅q(st+1,at+1;w)−q(st,at;w)\delta_t = y_t - q_t = r_t + \gamma \cdot q(s_{t+1}, a_{t+1}; \mathbf{w}) - q(s_t, a_t; \mathbf{w})δt=ytqt=rt+γq(st+1,at+1;w)q(st,at;w)

3.2 损失函数

Critic 的目标是让当前估计值接近 TD 目标,使用均方误差(MSE)损失:

L(w)=12[yt−q(st,at;w)]2=12δt2 L(\mathbf{w}) = \frac{1}{2} [y_t - q(s_t, a_t; \mathbf{w})]^2 = \frac{1}{2} \delta_t^2 L(w)=21[ytq(st,at;w)]2=21δt2

3.3 梯度计算

对损失函数求导:

∂L(w)∂w=∂∂w[12(yt−q(st,at;w))2]=(yt−q(st,at;w))⋅∂∂w[yt−q(st,at;w)]=(yt−q(st,at;w))⋅(−∂q(st,at;w)∂w)=−δt⋅∂q(st,at;w)∂w \begin{aligned} \frac{\partial L(\mathbf{w})}{\partial \mathbf{w}} &= \frac{\partial}{\partial \mathbf{w}} \left[ \frac{1}{2} (y_t - q(s_t, a_t; \mathbf{w}))^2 \right] \\ &= (y_t - q(s_t, a_t; \mathbf{w})) \cdot \frac{\partial}{\partial \mathbf{w}} [y_t - q(s_t, a_t; \mathbf{w})] \\ &= (y_t - q(s_t, a_t; \mathbf{w})) \cdot \left( -\frac{\partial q(s_t, a_t; \mathbf{w})}{\partial \mathbf{w}} \right) \\ &= -\delta_t \cdot \frac{\partial q(s_t, a_t; \mathbf{w})}{\partial \mathbf{w}} \end{aligned} wL(w)=w[21(ytq(st,at;w))2]=(ytq(st,at;w))w[ytq(st,at;w)]=(ytq(st,at;w))(wq(st,at;w))=δtwq(st,at;w)

注意:这里假设 TD 目标 yty_tyt 被视为常数(不对其求导),这是 TD 学习的标准做法。

3.4 Critic 更新规则

使用梯度下降法最小化损失:

w←w−αw⋅∂L(w)∂w=w−αw⋅(−δt⋅∂q(st,at;w)∂w) \begin{aligned} \mathbf{w} &\leftarrow \mathbf{w} - \alpha_w \cdot \frac{\partial L(\mathbf{w})}{\partial \mathbf{w}} \\ &= \mathbf{w} - \alpha_w \cdot \left( -\delta_t \cdot \frac{\partial q(s_t, a_t; \mathbf{w})}{\partial \mathbf{w}} \right) \end{aligned} wwαwwL(w)=wαw(δtwq(st,at;w))

最终得到:

w←w+αw⋅δt⋅∂q(st,at;w)∂w \boxed{\mathbf{w} \leftarrow \mathbf{w} + \alpha_w \cdot \delta_t \cdot \frac{\partial q(s_t, a_t; \mathbf{w})}{\partial \mathbf{w}}} ww+αwδtwq(st,at;w)

直观理解

  • δt>0\delta_t > 0δt>0 时,说明当前估计偏低,增大 q(st,at;w)q(s_t, a_t; \mathbf{w})q(st,at;w)
  • δt<0\delta_t < 0δt<0 时,说明当前估计偏高,减小 q(st,at;w)q(s_t, a_t; \mathbf{w})q(st,at;w)

4 算法流程

  1. 初始化:随机初始化 Actor 参数 θ\boldsymbol{\theta}θ 和 Critic 参数 w\mathbf{w}w
  2. 循环:对于每个时间步 ttt
    • 根据当前策略 π(a∣st;θ)\pi(a|s_t; \boldsymbol{\theta})π(ast;θ) 选择动作 ata_tat
    • 执行动作,观察奖励 rtr_trt 和下一状态 st+1s_{t+1}st+1
    • 选择下一动作 at+1∼π(⋅∣st+1;θ)a_{t+1} \sim \pi(\cdot|s_{t+1}; \boldsymbol{\theta})at+1π(st+1;θ)
    • 计算 TD 误差:δt=rt+γ⋅q(st+1,at+1;w)−q(st,at;w)\delta_t = r_t + \gamma \cdot q(s_{t+1}, a_{t+1}; \mathbf{w}) - q(s_t, a_t; \mathbf{w})δt=rt+γq(st+1,at+1;w)q(st,at;w)
    • 更新 Criticw←w+αw⋅δt⋅∂q(st,at;w)∂w\mathbf{w} \leftarrow \mathbf{w} + \alpha_w \cdot \delta_t \cdot \frac{\partial q(s_t, a_t; \mathbf{w})}{\partial \mathbf{w}}ww+αwδtwq(st,at;w)
    • 更新 Actorθ←θ+αθ⋅∂log⁡π(at∣st;θ)∂θ⋅q(st,at;w)\boldsymbol{\theta} \leftarrow \boldsymbol{\theta} + \alpha_{\theta} \cdot \frac{\partial \log \pi(a_t|s_t; \boldsymbol{\theta})}{\partial \boldsymbol{\theta}} \cdot q(s_t,a_t; \mathbf{w})θθ+αθθlogπ(atst;θ)q(st,at;w)
Logo

魔乐社区(Modelers.cn) 是一个中立、公益的人工智能社区,提供人工智能工具、模型、数据的托管、展示与应用协同服务,为人工智能开发及爱好者搭建开放的学习交流平台。社区通过理事会方式运作,由全产业链共同建设、共同运营、共同享有,推动国产AI生态繁荣发展。

更多推荐