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()