Files
RL-Study/TRPO/agent.py
T

296 lines
13 KiB
Python
Raw Normal View History

import torch
import torch.nn as nn
import numpy as np
from models import ActorNet, CriticNet
import torch.optim as optim
from torch.distributions import Normal
# ==========================================
# 辅助工具函数:参数与向量的互相转换
# ==========================================
def get_flat_params_from(model):
"""
把模型中散落在各个层的所有参数,按顺序拼接成一个巨大的一维 Tensor。
相当于数学推导中的参数向量 theta。
"""
params = []
for param in model.parameters():
params.append(param.data.view(-1))
return torch.cat(params)
def set_flat_params_to(model, flat_params):
"""
把计算好的新一维参数向量,按照对应尺寸还原、塞回神经网络的各个层中。
这是用来真正执行 theta_new 赋值的。
"""
prev_ind = 0
for param in model.parameters():
flat_size = int(np.prod(param.size()))
# 截取对应长度的数据,并 reshape 回原始层的形状
param.data.copy_(flat_params[prev_ind:prev_ind + flat_size].view(param.size()))
prev_ind += flat_size
def get_flat_grad_from(loss, model):
"""
对指定的 loss 求网络参数的梯度,并直接压平成一个一维 Tensor 返回。
相当于计算梯度向量 g。
"""
# retain_graph=True 是因为我们后面算二阶导可能还会用到当前的计算图
grads = torch.autograd.grad(loss, model.parameters(), retain_graph=True)
return torch.cat([grad.view(-1) for grad in grads])
# ==========================================
# TRPO 核心数学引擎
# ==========================================
def conjugate_gradient(fvp_func, b, nsteps=10, residual_tol=1e-10):
"""
共轭梯度法 (Conjugate Gradient, CG)
用于近似求解线性方程组 Ax = b,在这里也就是求解 As = g。
它不需要矩阵 A,只需要一个能计算 A*v 的函数 (也就是下面的 fvp_func)。
参数:
fvp_func: 传入一个向量 v,返回 Fisher矩阵乘该向量的结果 A*v
b: 目标向量 (在 TRPO 中就是策略梯度向量 g)
nsteps: 迭代次数 (通常 10 次就能逼近得很好)
"""
x = torch.zeros_like(b) # 初始解设为 0
r = b.clone() # 初始残差
p = b.clone() # 初始搜索方向
rdotr = torch.dot(r, r) # 残差的内积
for i in range(nsteps):
# 计算 A * p (也就是 FVP)
Ap = fvp_func(p)
# 计算步长 alpha
alpha = rdotr / (torch.dot(p, Ap) + 1e-8)
# 更新解 x
x += alpha * p
# 更新残差 r
r -= alpha * Ap
new_rdotr = torch.dot(r, r)
# 如果残差已经足够小,提前退出
if new_rdotr < residual_tol:
break
# 计算方向更新系数 beta
beta = new_rdotr / rdotr
# 更新搜索方向 p
p = r + beta * p
rdotr = new_rdotr
return x
def fisher_vector_product(actor_net, states, vector, damping=0.1):
"""
海森向量积 / 费雪信息矩阵-向量积 (FVP)
这是 TRPO 的绝对核心黑科技:通过连续两次自动求导,计算 A * v,无需显式构造 A!
"""
# 1. 用当前的策略网络计算出旧的均值和标准差 (停止梯度更新,作为基准点)
mean_old, std_old = actor_net(states)
mean_old = mean_old.detach()
std_old = std_old.detach()
# 2. 重新进行一次前向传播,保留计算图
mean, std = actor_net(states)
# 3. 解析计算 KL 散度 (高斯分布的精确闭式解)
# 公式: log(std/std_old) + (std_old^2 + (mean_old - mean)^2) / (2 * std^2) - 0.5
# 注意:TRPO 是对状态空间求期望,所以最后要求均值 (mean)
kl = torch.log(std / std_old) + (std_old.pow(2) + (mean_old - mean).pow(2)) / (2.0 * std.pow(2)) - 0.5
kl = kl.sum(dim=1, keepdim=True).mean()
# 4. 第一次求导:计算 KL 对网络参数的一阶梯度 (Jacobian)
# create_graph=True 极其关键:它让一阶梯度本身也成为计算图的一部分,为求二阶导做准备
grads = torch.autograd.grad(kl, actor_net.parameters(), create_graph=True)
flat_grad_kl = torch.cat([grad.view(-1) for grad in grads])
# 5. 计算一阶梯度向量与传入向量 v 的内积
# 这个点积的结果是一个标量 (Scalar)
kl_v = torch.dot(flat_grad_kl, vector)
# 6. 第二次求导:对上面的点积标量再次求参数的梯度
# 根据微积分法则,梯度的点积的梯度 = Hessian * v
grads_v = torch.autograd.grad(kl_v, actor_net.parameters())
flat_grad_grad_kl = torch.cat([grad.contiguous().view(-1) for grad in grads_v]).detach()
# 7. 加上阻尼项 (Damping)
# 给对角线加上一个小常数 (damping * vector),确保 FIM 矩阵正定,提升 CG 求解的数值稳定性
return flat_grad_grad_kl + vector * damping
class TRPOAgent:
def __init__(self, state_dim, action_dim, max_kl=0.01, cg_iters=10, cg_residual_tol=1e-10, cg_damping=0.1):
"""
初始化 TRPO 智能体
"""
self.actor = ActorNet(state_dim, action_dim)
self.critic = CriticNet(state_dim)
# Critic 使用普通的 Adam 优化器即可 (学习率设为 1e-3)
self.critic_optimizer = optim.Adam(self.critic.parameters(), lr=1e-3)
# TRPO 超参数
self.max_kl = max_kl # 也就是公式里的 delta (信任域边界)
self.cg_iters = cg_iters # 共轭梯度法的迭代次数
self.cg_residual_tol = cg_residual_tol
self.cg_damping = cg_damping # 海森矩阵的阻尼系数
def compute_surrogate_obj(self, states, actions, old_log_probs, advantages):
"""
计算替代目标函数 (Surrogate Objective)
公式: L = E[ (pi_new / pi_old) * A ]
"""
# 计算当前策略下动作的对数概率
mean, std = self.actor(states)
dist = Normal(mean, std)
# 注意: 如果动作是多维的,需要对各维度的 log_prob 求和
new_log_probs = dist.log_prob(actions).sum(dim=1, keepdim=True)
# 计算重要性采样比率 (Ratio): exp(log_new - log_old) = new / old
ratio = torch.exp(new_log_probs - old_log_probs)
# 替代目标函数 (最大化目标,所以返回均值)
surrogate_obj = (ratio * advantages).mean()
return surrogate_obj
def update(self, rollout_buffer, next_state, done):
"""
执行一次完整的 TRPO 更新 (Actor 和 Critic)
参数:
rollout_buffer: 收集好数据的经验池
next_state: 轨迹结束时的下一个状态 (用于计算 GAE)
done: 轨迹是否结束的标志位
"""
# 在更新前,先用 Critic 算出 GAE 和 Returns
with torch.no_grad():
# 如果环境已经 done(比如倒立摆摔倒或超时),那未来的预期价值就是 0
# 否则,用 Critic 网络预测一下 next_state 的价值
if done:
last_value = 0.0
else:
next_state_tensor = torch.tensor(next_state, dtype=torch.float32)
last_value = self.critic(next_state_tensor).item()
# 真正调用 buffer 的函数,计算出 numpy 格式的 returns 和 advantages
returns_np, advantages_np = rollout_buffer.compute_returns_and_advantages(last_value)
# 转成深度学习需要的 Tensor,并对齐维度 [Batch, 1]
returns = torch.tensor(returns_np, dtype=torch.float32).view(-1, 1)
advantages = torch.tensor(advantages_np, dtype=torch.float32).view(-1, 1)
# 1. 从经验池中获取并整理数据
states, actions, old_log_probs = rollout_buffer.get_data()
# ==========================================
# 第一步:计算目标函数的一阶梯度 (g)
# ==========================================
surrogate_obj_old = self.compute_surrogate_obj(states, actions, old_log_probs, advantages)
# 注意 TRPO 是最大化目标,所以 loss 是负的 objective
loss_actor = -surrogate_obj_old
# 使用我们在 Stage 3 写的辅助函数获取铺平的一阶梯度 g
g = get_flat_grad_from(loss_actor, self.actor).detach()
# ==========================================
# 第二步:用共轭梯度法 (CG) 求自然梯度方向 (s)
# ==========================================
# 定义一个局部函数,把 states 封进去,专供 CG 调用
def fvp_callable(v):
return fisher_vector_product(self.actor, states, v, self.cg_damping)
# 解方程 As = g,得到搜索方向 step_dir (即公式里的 s)
step_dir = conjugate_gradient(fvp_callable, -g, self.cg_iters, self.cg_residual_tol)
# ==========================================
# 第三步:计算理论最大步长 (beta)
# ==========================================
# s^T A s (通过 FVP 再算一次)
sAs = torch.dot(step_dir, fvp_callable(step_dir))
# 为了防止除以 0 的数值不稳定,加上 1e-8
# 公式: beta = sqrt( 2 * delta / (s^T A s) )
beta = torch.sqrt(2 * self.max_kl / (sAs + 1e-8))
# 完整的最大更新步长向量
full_step = beta * step_dir
# ==========================================
# 第四步:回溯线搜索 (Backtracking Line Search)
# ==========================================
# 获取当前 Actor 的初始参数
old_params = get_flat_params_from(self.actor)
success = False
fraction = 1.0 # 步长衰减系数的初始值
# 尝试 10 次缩小步长
for i in range(10):
# 试探性地迈出一步: theta_new = theta_old + fraction * full_step
new_params = old_params + fraction * full_step
set_flat_params_to(self.actor, new_params)
# 在新参数下,重新评估目标函数和 KL 散度
with torch.no_grad():
# 1. 检查目标函数是否真的提升了
surrogate_obj_new = self.compute_surrogate_obj(states, actions, old_log_probs, advantages)
improvement = surrogate_obj_new - surrogate_obj_old
# 2. 检查 KL 散度是否满足约束 (<= max_kl)
mean_new, std_new = self.actor(states)
mean_old, std_old = self.actor(states) # 注意这里应该用一开始记录的固定不变的 old 值,为了严谨,我们在最开始算一次并脱离计算图
# (为了简化代码,更严谨的做法是在线搜索外先算好 old 分布参数传进来)
# 我们在这里写一个快速的 KL 检查逻辑
mean_old_frozen, std_old_frozen = self.actor(states)
set_flat_params_to(self.actor, old_params) # 临时切回老参数获取冻结的分布
mean_old_frozen, std_old_frozen = self.actor(states)
mean_old_frozen, std_old_frozen = mean_old_frozen.detach(), std_old_frozen.detach()
set_flat_params_to(self.actor, new_params) # 再切回新参数算 KL
mean_new, std_new = self.actor(states)
kl = torch.log(std_new / std_old_frozen) + (std_old_frozen.pow(2) + (mean_old_frozen - mean_new).pow(2)) / (2.0 * std_new.pow(2)) - 0.5
kl_mean = kl.sum(dim=1).mean()
# 判断双重安全条件:目标提升了,且 KL 没超标
if improvement.item() > 0 and kl_mean.item() <= self.max_kl:
success = True
print(f"线搜索成功: 迭代次数 {i}, 提升值 {improvement.item():.4f}, KL 散度 {kl_mean.item():.4f}")
break
else:
# 如果失败了,步长减半,再试一次
fraction *= 0.5
# 如果 10 次尝试都失败了,说明这里地形太差,安全起见我们不更新 Actor 了
if not success:
print("线搜索失败,放弃此次 Actor 更新,保持原参数。")
set_flat_params_to(self.actor, old_params)
# ==========================================
# 第五步:更新 Critic (价值网络)
# ==========================================
# Critic 的更新非常简单,就是标准的深度学习监督训练,目标是逼近真实的 Return
criterion = nn.MSELoss()
# 多次迭代更新 Critic 以充分拟合价值 (通常设为 10 次)
for _ in range(10):
values = self.critic(states)
loss_critic = criterion(values, returns)
self.critic_optimizer.zero_grad()
loss_critic.backward()
self.critic_optimizer.step()
# 清空经验池,为下一次环境交互做准备
rollout_buffer.clear()