策略梯度方法里有个绕不开的老问题:一轮更新到底应该走多大。步长太小,训练慢;步长太大,一个 batch 的噪声就可能把策略推到悬崖边,收益曲线瞬间崩掉。TRPO(Trust Region Policy Optimization,信赖域策略优化)就是为解决这个步长困境提出的算法,它用 KL 散度把每次更新限制在一个“可信的区域”内,从而在理论上保证策略改进的单调性。也正因为 TRPO 给出了稳定更新的约束思路,后来的 PPO 才可以直接用裁剪目标函数来近似同样的效果,所以 TRPO 通常被称为 PPO 的前身。
许多人在学习强化学习时,先接触 DQN,再接触策略梯度,然后直接跳到 PPO,最后回头补 TRPO。这种顺序虽然能尽快用上 PPO,但容易留下一个模糊地带:PPO 的 clip 为什么能把策略限制住?它和理论上的“信任域”有什么关系?要回答这些问题,必须回到 TRPO。这篇文章围绕 TRPO 的核心思路、数学原理、算法流程、实现细节和与 PPO 的对比展开,适合已经了解策略梯度基本公式、想深入理解 RL 优化原理的读者。
1. TRPO 出现之前,策略梯度为什么难调
1.1 策略梯度的标准形式和方差来源
策略梯度方法(Policy Gradient)直接对策略参数 θ 求期望回报的梯度。目标是最大化:
J(θ) = E_{τ ~ π_θ} [ Σ_{t=0}^{T} γ^t r_t ]
用经典的 REINFORCE 或带基线的策略梯度公式表示:
∇_θ J(θ) = E_{τ, s_t, a_t} [ ∇_θ log π_θ(a_t | s_t) · A_t ]
其中 A_t 是优势函数,表示当前动作相对平均水平的优势。这个公式看起来简洁,但实际落地时有两个问题。
第一个问题是方差。回报是由一条完整轨迹汇总出来的,而轨迹中每一步都包含随机因素。同一个策略在相同环境下采样两次,得到的回报可能差很多。用一批有限样本估计梯度,噪声会被放大。第二个问题是更新方向对步长非常敏感。策略梯度给出的是当前策略下的局部方向,离开当前参数点之后,这个方向就不再可靠。如果用固定学习率大步更新,策略分布可能剧烈偏移,下一轮采样到的数据质量下降,训练曲线直接崩坏。
1.2 固定步长为什么会让训练崩溃
常见的策略梯度实现使用 Adam 或 SGD 更新参数:
θ_{k+1} = θ_k + α ∇_θ J(θ_k)
学习率 α 设置得小,训练会慢,但不至于立刻崩。然而在许多连续控制问题里,策略输出的是一个高斯分布,均值变化一点点,采样出来的动作分布可能完全不同。一轮更新后,如果新策略和旧策略的重叠区域太小,旧数据评估出来的目标函数就不再准确,下一轮更新就建立在一个错误的方向上。
更麻烦的是,策略网络的损失函数和普通监督学习不一样。监督学习的标签是固定的,模型的输出变化不会改变数据分布;强化学习的损失依赖当前策略采样出的数据,策略一变,数据分布就变。这就是“非平稳目标”问题。固定学习率没有考虑新旧策略之间的距离,所以容易出现“更新一次、性能骤降、再也回不来”的情况。
1.3 TRPO 的解决思路:给更新画一条安全边界
TRPO 的出发点是:不要只看梯度方向走得多远,而是在每次更新前,先检查新策略和旧策略的 KL 散度是否超过阈值。如果超过,就缩短步长,直到 KL 散度回到安全范围内。
这里的 KL 散度不是衡量参数距离,而是衡量策略分布之间的距离:
D_KL(π_{θ_old} ‖ π_θ) = E_{a ~ π_{θ_old}} [ log π_{θ_old}(a | s) - log π_θ(a | s) ]
当新旧策略分布差异较大时,KL 散度也变大。TRPO 每次更新都要求这个值不超过 δ,例如 δ = 0.01。这样策略不会一次性偏离太远,训练过程在理论上具备单调改进保证。这个“不超过 δ”的约束,就是信赖域(Trust Region)的含义。
2. TRPO 的数学原理:目标函数、KL 约束和单调改进
2.1 surrogate 目标函数怎么构造
TRPO 并不直接优化真实期望回报,因为真实期望回报无法对不同 θ 直接计算。它使用重要性采样(Importance Sampling)构造一个代理目标函数:
L(θ) = E_{(s, a) ~ π_{θ_old}} [ (π_θ(a | s) / π_{θ_old}(a | s)) · A(s, a) ]
其中 π_θ(a | s) / π_{θ_old}(a | s) 是重要性采样比率。旧策略采样得到的轨迹,仍然可以用来估计新策略的目标函数,只要新旧策略差异不大。写成对数形式:
r_t(θ) = exp( log π_θ(a_t | s_t) - log π_{θ_old}(a_t | s_t) )
L(θ) = E_t [ r_t(θ) · A_t ]
这里有一个重要细节:log π_{θ_old}(a_t | s_t) 在更新时必须是固定值,不能参与梯度计算,否则比率会被错误地“自我放大”。
2.2 KL 散度约束为什么比惩罚项稳
一种自然想法是把 KL 散度作为惩罚项加到目标函数里:
maximize L(θ) - β · D_KL(π_{θ_old} ‖ π_θ)
这样也能限制步长,但 β 很难调。β 太小,约束失去作用;β 太大,策略几乎不动。TRPO 选择把 KL 散度当作硬约束,而不是惩罚项:
maximize L(θ) subject to D_KL(π_{θ_old} ‖ π_θ) ≤ δ
这个选择来自一个理论推导。TRPO 论文参考了 Kakade 和 Langford 的保守策略迭代工作,给出了真实回报与 surrogate 目标之间的下界关系。满足 KL 约束时,策略可以保证在某一置信水平内单调改进。用惩罚项时,这个理论保证很难直接成立;用硬约束时,至少每一次更新都有明确的“安全边界”。
实际使用中,TRPO 通常使用平均 KL 散度约束,而不是对每个状态都施加最大 KL 约束。平均 KL 更容易估计,计算开销低;理论上最大 KL 更严格,但实现更复杂。论文中使用平均 KL 约束也取得了稳定效果。阅读代码时要注意区分这两种写法。
2.3 从理论下界到实际算法
TRPO 的理论保证核心是一个下界表达式:
η(π_θ) ≥ L(θ) - C · max_s D_KL(π_{θ_old}(· | s) ‖ π_θ(· | s))
其中 η 是真实期望回报,C 是由折扣因子和奖励范围决定的常数。这个式子说明,只要新旧策略的 KL 散度足够小,真实回报就能被 surrogate 目标近似,更新方向就是可信的。TRPO 的实际做法是把 max 的 KL 换成平均 KL,从而让问题变得可求解。
由于直接求解带约束的深度网络优化非常困难,TRPO 没有把 KL 约束直接丢给通用优化器,而是拆成两个阶段:先用二阶信息求出理论最优步长方向,再通过线搜索保证 KL 约束被满足。这正是下一节要展开的内容。
3. TRPO 算法完整流程
3.1 采样与 Advantage 估计
TRPO 的每一步更新都遵循“采样-估计-更新”的循环。先从当前策略 π_{θ_old} 中采样一批轨迹,计算每条轨迹上每个时间步的回报。为了降低方差,通常使用 GAE(Generalized Advantage Estimation)计算优势:
A_t = Σ_{l=0}^{∞} (γλ)^l δ_{t+l} δ_t = r_t + γ V(s_{t+1}) - V(s_t)
GAE 有两个超参数:γ 控制折扣幅度,λ 控制偏差和方差的权衡。λ 越大,方差越大但偏差越小;λ 越小,估计越平滑但偏差越大。TRPO 实践里常见 γ 取 0.99,λ 取 0.95,但最终取值要看具体任务。
需要强调:TRPO 需要一条单独的 Critic 网络来估计价值函数 V(s),因为优势估计依赖价值函数。Critic 可以用回归损失更新,策略网络则使用 TRPO 的约束优化更新。
3.2 共轭梯度法如何避开逆矩阵
如果直接对带约束的目标函数做二阶优化,需要计算 Fisher 信息矩阵 F 的逆:
θ_{new} = θ_{old} + α F^{-1} ∇L(θ_{old})
Fisher 信息矩阵的大小是参数总数 × 参数总数。一个只有 10 万参数的策略网络,Fisher 矩阵就是 10 万 × 10 万,存储和求逆都不现实。TRPO 采用共轭梯度法(Conjugate Gradient)求解线性方程 F x = g,不需要显式构造 F,只需要能够计算 F 与任意向量 v 的乘积 F v。
F v 可以通过 KL 散度的 Hessian-vector product 来近似。令:
KL = E_s [ D_KL(π_{θ_old}(· | s) ‖ π_θ(· | s)) ]
对 KL 求一阶梯度得到 g_kl,再计算 g_kl 和 v 的点积对参数的二阶梯度,得到 F v。这个操作的复杂度与一次反向传播接近,因此可以接受。
共轭梯度法的迭代次数一般取 10 到 20 次。每一次都只做向量乘积,不需要存储矩阵,所以内存可控。求解完成后得到的 x 就是近似的自然梯度方向。
3.3 线搜索与最终更新
共轭梯度给出的是方向 x,但步长还不能随意设置。TRPO 先根据 KL 约束计算最大步长:
step_size = sqrt(2δ / (x^T F x))
然后从完整步长开始尝试,一步一步缩小。每次尝试都设置新参数,重新计算 KL 散度和 surrogate loss。如果 KL 超过 δ,或者 surrogate loss 没有改善,就退回一半步长继续尝试。这个过程叫线搜索(Line Search)。
最终更新公式是:
θ_{new} = θ_{old} + step_size · x
TRPO 的稳定性来自这道双保险:共轭梯度给出合理的二阶更新方向,线搜索负责验证每一步更新确实停留在可信域内。即使理论上计算出来的步长偏大,线搜索也能兜底。
4. TRPO 核心模块实现(PyTorch 风格)
4.1 参数扁平化与工具函数
实现 TRPO 时,第一个坑是参数形态。神经网络参数是嵌套的 Tensor,而共轭梯度要求把参数当作一维向量处理。需要把参数拍平、恢复、再拍平。
import torch import torch.nn as nn from torch.distributions import Independent, Normal def get_flat_params(model): return torch.cat([p.data.view(-1) for p in model.parameters()]) def set_flat_params(model, flat_params): idx = 0 for p in model.parameters(): n = p.numel() p.data.copy_(flat_params[idx:idx + n].view(p.shape)) idx += n def flat_grad(f, params, retain_graph=True, create_graph=False): grads = torch.autograd.grad(f, params, retain_graph=retain_graph, create_graph=create_graph) return torch.cat([g.contiguous().view(-1) for g in grads])这里flat_grad的create_graph参数很关键。计算目标函数梯度时不需要二阶信息,设为 False;计算 KL 的一阶梯度并继续求二阶时,需要设为 True,否则后面无法对参数再次求导。
4.2 高斯策略和 KL 估计
下面是一个简单的连续动作策略网络,输出高斯分布的均值和对数标准差。
class GaussianPolicy(nn.Module): def __init__(self, state_dim, action_dim, hidden=64): super().__init__() self.net = nn.Sequential( nn.Linear(state_dim, hidden), nn.Tanh(), nn.Linear(hidden, hidden), nn.Tanh(), ) self.mean_head = nn.Linear(hidden, action_dim) self.logstd = nn.Parameter(torch.zeros(action_dim)) def forward(self, obs): mean = self.mean_head(self.net(obs)) std = torch.exp(self.logstd) dist = Independent(Normal(mean, std), 1) return dist使用时,策略返回的是一个 PyTorch 分布对象,可以直接调用log_prob和sample。在 TRPO 中,旧策略的分布需要固定,因此更新前要对旧策略做一次完整复制,并关闭梯度:
old_policy = GaussianPolicy(state_dim, action_dim) old_policy.load_state_dict(policy