Files

296 lines
13 KiB
Python

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