240 lines
9.5 KiB
Python
240 lines
9.5 KiB
Python
import numpy as np
|
|||
|
|
import torch
|
||
|
|
import torch.nn.functional as F
|
||
|
|
from torch.distributions import Normal
|
||
|
|
from torch.nn.utils import parameters_to_vector, vector_to_parameters
|
||
|
|
from networks import PolicyNetwork, ValueNetwork
|
||
|
|
|
||
|
|
class TRPOAgent:
|
||
|
|
def __init__(self, state_dim, action_dim, action_bound, hidden_dim=128, kl_margin=0.01, gamma=0.99, tau=0.97, cg_iters=10):
|
||
|
|
self.gamma = gamma
|
||
|
|
self.tau = tau
|
||
|
|
self.kl_margin = kl_margin
|
||
|
|
self.cg_iters = cg_iters
|
||
|
|
|
||
|
|
self.actor = PolicyNetwork(state_dim, action_dim, action_bound, hidden_dim)
|
||
|
|
self.critic = ValueNetwork(state_dim, hidden_dim)
|
||
|
|
self.critic_optimizer = torch.optim.Adam(self.critic.parameters(), lr=1e-3)
|
||
|
|
|
||
|
|
def get_action(self, state):
|
||
|
|
"""
|
||
|
|
根据当前状态采样动作
|
||
|
|
"""
|
||
|
|
state_tensor = torch.FloatTensor(state).unsqueeze(0)
|
||
|
|
with torch.no_grad():
|
||
|
|
dist = self.actor.evaluate(state_tensor)
|
||
|
|
action = dist.sample()
|
||
|
|
return action.squeeze(0).numpy()
|
||
|
|
|
||
|
|
def _compute_advantages(self, rewards, values, masks):
|
||
|
|
"""
|
||
|
|
使用广义优势估计 (GAE) 计算优势函数
|
||
|
|
"""
|
||
|
|
returns = []
|
||
|
|
gae = 0
|
||
|
|
for i in reversed(range(len(rewards))):
|
||
|
|
delta = rewards[i] + self.gamma * values[i + 1] * masks[i] - values[i]
|
||
|
|
gae = delta + self.gamma * self.tau * masks[i] * gae
|
||
|
|
returns.insert(0, gae + values[i])
|
||
|
|
return returns
|
||
|
|
|
||
|
|
# --- 关键修复 1:将固定的 old_dist 作为参数传入 ---
|
||
|
|
def _hessian_vector_product(self, states, old_dist, vector, damping=0.1):
|
||
|
|
"""
|
||
|
|
计算海森矩阵与向量的乘积 (Hvp)
|
||
|
|
这里的 old_dist 必须是不带梯度的历史常数分布
|
||
|
|
"""
|
||
|
|
dist = self.actor.evaluate(states)
|
||
|
|
# 计算当前分布与固定的旧分布之间的 KL 散度
|
||
|
|
kl = torch.distributions.kl_divergence(old_dist, dist).mean()
|
||
|
|
|
||
|
|
# 一阶导数
|
||
|
|
grads = torch.autograd.grad(kl, self.actor.parameters(), create_graph=True)
|
||
|
|
flat_grad_kl = torch.cat([grad.view(-1) for grad in grads])
|
||
|
|
|
||
|
|
# 与给定向量点乘
|
||
|
|
kl_v = (flat_grad_kl * vector).sum()
|
||
|
|
|
||
|
|
# 二阶导数
|
||
|
|
grads = torch.autograd.grad(kl_v, self.actor.parameters())
|
||
|
|
flat_grad_grad_kl = torch.cat([grad.contiguous().view(-1) for grad in grads])
|
||
|
|
|
||
|
|
# 加入阻尼系数,保证矩阵正定,避免数值不稳定发散
|
||
|
|
return flat_grad_grad_kl + vector * damping
|
||
|
|
|
||
|
|
# --- 关键修复 2:共轭梯度法同样接收 old_dist ---
|
||
|
|
def _conjugate_gradient(self, states, old_dist, b, nsteps, residual_tol=1e-10):
|
||
|
|
"""
|
||
|
|
共轭梯度法,近似求解 Hx = b
|
||
|
|
"""
|
||
|
|
x = torch.zeros_like(b)
|
||
|
|
r = b.clone()
|
||
|
|
p = b.clone()
|
||
|
|
rdotr = torch.dot(r, r)
|
||
|
|
for _ in range(nsteps):
|
||
|
|
Hp = self._hessian_vector_product(states, old_dist, p)
|
||
|
|
alpha = rdotr / torch.dot(p, Hp)
|
||
|
|
x += alpha * p
|
||
|
|
r -= alpha * Hp
|
||
|
|
new_rdotr = torch.dot(r, r)
|
||
|
|
if new_rdotr < residual_tol:
|
||
|
|
break
|
||
|
|
p = r + new_rdotr / rdotr * p
|
||
|
|
rdotr = new_rdotr
|
||
|
|
return x
|
||
|
|
|
||
|
|
def update(self, memory):
|
||
|
|
states = torch.FloatTensor(np.array([m[0] for m in memory]))
|
||
|
|
actions = torch.FloatTensor(np.array([m[1] for m in memory]))
|
||
|
|
rewards = [m[2] for m in memory]
|
||
|
|
next_states = torch.FloatTensor(np.array([m[3] for m in memory]))
|
||
|
|
masks = [m[4] for m in memory]
|
||
|
|
|
||
|
|
with torch.no_grad():
|
||
|
|
values = self.critic(states).squeeze().numpy().tolist()
|
||
|
|
next_value = self.critic(next_states[-1].unsqueeze(0)).squeeze().item()
|
||
|
|
values.append(next_value)
|
||
|
|
|
||
|
|
returns = self._compute_advantages(rewards, values, masks)
|
||
|
|
returns = torch.FloatTensor(returns)
|
||
|
|
values = torch.FloatTensor(values[:-1])
|
||
|
|
advantages = returns - values
|
||
|
|
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
|
||
|
|
|
||
|
|
# 【修复】:加大 Critic 的训练力度,从 10 提升到 40 Epochs
|
||
|
|
# 确保裁判的眼光足够准确
|
||
|
|
for _ in range(40):
|
||
|
|
critic_loss = F.mse_loss(self.critic(states).squeeze(), returns)
|
||
|
|
self.critic_optimizer.zero_grad()
|
||
|
|
critic_loss.backward()
|
||
|
|
self.critic_optimizer.step()
|
||
|
|
|
||
|
|
with torch.no_grad():
|
||
|
|
old_mean, old_std = self.actor(states)
|
||
|
|
old_dist = Normal(old_mean, old_std)
|
||
|
|
old_log_probs = old_dist.log_prob(actions).sum(dim=1)
|
||
|
|
|
||
|
|
def compute_surrogate_loss():
|
||
|
|
dist = self.actor.evaluate(states)
|
||
|
|
log_probs = dist.log_prob(actions).sum(dim=1)
|
||
|
|
ratio = torch.exp(log_probs - old_log_probs)
|
||
|
|
surrogate_loss = (ratio * advantages).mean()
|
||
|
|
return surrogate_loss, dist
|
||
|
|
|
||
|
|
surrogate_loss, dist = compute_surrogate_loss()
|
||
|
|
|
||
|
|
loss_grad = torch.autograd.grad(surrogate_loss, self.actor.parameters())
|
||
|
|
loss_grad_flat = torch.cat([grad.view(-1) for grad in loss_grad])
|
||
|
|
|
||
|
|
step_dir = self._conjugate_gradient(states, old_dist, loss_grad_flat, self.cg_iters)
|
||
|
|
|
||
|
|
shs = 0.5 * torch.dot(step_dir, self._hessian_vector_product(states, old_dist, step_dir))
|
||
|
|
|
||
|
|
if shs < 1e-8:
|
||
|
|
return
|
||
|
|
|
||
|
|
lm = torch.sqrt(shs / self.kl_margin)
|
||
|
|
fullstep = step_dir / lm
|
||
|
|
|
||
|
|
old_params = parameters_to_vector(self.actor.parameters())
|
||
|
|
|
||
|
|
# 线性搜索
|
||
|
|
success = False
|
||
|
|
step_size = 1.0
|
||
|
|
for _ in range(10):
|
||
|
|
new_params = old_params + step_size * fullstep
|
||
|
|
vector_to_parameters(new_params, self.actor.parameters())
|
||
|
|
|
||
|
|
with torch.no_grad():
|
||
|
|
new_surrogate_loss, new_dist = compute_surrogate_loss()
|
||
|
|
kl = torch.distributions.kl_divergence(old_dist, new_dist).mean()
|
||
|
|
|
||
|
|
# 【修复】:增加极小的浮点数宽容度,防止在极小提升时被误判失败而拒绝更新
|
||
|
|
if new_surrogate_loss >= surrogate_loss - 1e-8 and kl <= self.kl_margin * 1.5:
|
||
|
|
success = True
|
||
|
|
break
|
||
|
|
|
||
|
|
step_size *= 0.5
|
||
|
|
|
||
|
|
if not success:
|
||
|
|
vector_to_parameters(old_params, self.actor.parameters())
|
||
|
|
"""
|
||
|
|
利用收集到的轨迹数据更新 Actor 和 Critic
|
||
|
|
"""
|
||
|
|
states = torch.FloatTensor(np.array([m[0] for m in memory]))
|
||
|
|
actions = torch.FloatTensor(np.array([m[1] for m in memory]))
|
||
|
|
rewards = [m[2] for m in memory]
|
||
|
|
next_states = torch.FloatTensor(np.array([m[3] for m in memory]))
|
||
|
|
masks = [m[4] for m in memory]
|
||
|
|
|
||
|
|
# 1. 拟合价值网络 (Critic)
|
||
|
|
with torch.no_grad():
|
||
|
|
values = self.critic(states).squeeze().numpy().tolist()
|
||
|
|
next_value = self.critic(next_states[-1].unsqueeze(0)).squeeze().item()
|
||
|
|
values.append(next_value)
|
||
|
|
|
||
|
|
returns = self._compute_advantages(rewards, values, masks)
|
||
|
|
returns = torch.FloatTensor(returns)
|
||
|
|
values = torch.FloatTensor(values[:-1])
|
||
|
|
advantages = returns - values
|
||
|
|
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
|
||
|
|
|
||
|
|
for _ in range(10):
|
||
|
|
critic_loss = F.mse_loss(self.critic(states).squeeze(), returns)
|
||
|
|
self.critic_optimizer.zero_grad()
|
||
|
|
critic_loss.backward()
|
||
|
|
self.critic_optimizer.step()
|
||
|
|
|
||
|
|
# --- 关键修复 3:在截断梯度的环境下生成严格的旧分布 ---
|
||
|
|
with torch.no_grad():
|
||
|
|
old_mean, old_std = self.actor(states)
|
||
|
|
old_dist = Normal(old_mean, old_std)
|
||
|
|
old_log_probs = old_dist.log_prob(actions).sum(dim=1)
|
||
|
|
|
||
|
|
def compute_surrogate_loss():
|
||
|
|
# 计算替代目标函数 (Surrogate Objective)
|
||
|
|
dist = self.actor.evaluate(states)
|
||
|
|
log_probs = dist.log_prob(actions).sum(dim=1)
|
||
|
|
ratio = torch.exp(log_probs - old_log_probs)
|
||
|
|
surrogate_loss = (ratio * advantages).mean()
|
||
|
|
return surrogate_loss, dist
|
||
|
|
|
||
|
|
surrogate_loss, dist = compute_surrogate_loss()
|
||
|
|
|
||
|
|
loss_grad = torch.autograd.grad(surrogate_loss, self.actor.parameters())
|
||
|
|
loss_grad_flat = torch.cat([grad.view(-1) for grad in loss_grad])
|
||
|
|
|
||
|
|
# 传入 old_dist,确保海森矩阵计算包含准确的曲率信息
|
||
|
|
step_dir = self._conjugate_gradient(states, old_dist, loss_grad_flat, self.cg_iters)
|
||
|
|
|
||
|
|
shs = 0.5 * torch.dot(step_dir, self._hessian_vector_product(states, old_dist, step_dir))
|
||
|
|
|
||
|
|
# 增加数值保护:防止 shs 出现负数或极小值导致报错
|
||
|
|
if shs < 1e-8:
|
||
|
|
return
|
||
|
|
|
||
|
|
lm = torch.sqrt(shs / self.kl_margin)
|
||
|
|
fullstep = step_dir / lm
|
||
|
|
|
||
|
|
old_params = parameters_to_vector(self.actor.parameters())
|
||
|
|
|
||
|
|
# 线性搜索 (Line Search)
|
||
|
|
success = False
|
||
|
|
step_size = 1.0
|
||
|
|
for _ in range(10):
|
||
|
|
new_params = old_params + step_size * fullstep
|
||
|
|
vector_to_parameters(new_params, self.actor.parameters())
|
||
|
|
|
||
|
|
with torch.no_grad():
|
||
|
|
new_surrogate_loss, new_dist = compute_surrogate_loss()
|
||
|
|
kl = torch.distributions.kl_divergence(old_dist, new_dist).mean()
|
||
|
|
|
||
|
|
if new_surrogate_loss > surrogate_loss and kl <= self.kl_margin:
|
||
|
|
success = True
|
||
|
|
break
|
||
|
|
|
||
|
|
step_size *= 0.5
|
||
|
|
|
||
|
|
if not success:
|
||
|
|
vector_to_parameters(old_params, self.actor.parameters())
|