Files
RL_TRPO/agent/trpo.py
T

222 lines
8.3 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.95,
cg_iters=10,
critic_epochs=40,
device="cpu",
):
self.gamma = gamma
self.tau = tau
self.kl_margin = kl_margin
self.cg_iters = cg_iters
self.critic_epochs = critic_epochs
self.device = torch.device(device)
self.actor = PolicyNetwork(state_dim, action_dim, action_bound, hidden_dim).to(self.device)
self.critic = ValueNetwork(state_dim, hidden_dim).to(self.device)
self.critic_optimizer = torch.optim.Adam(self.critic.parameters(), lr=1e-3)
def _to_tensor(self, array_like):
return torch.as_tensor(array_like, dtype=torch.float32, device=self.device)
def get_action(self, state):
"""
根据当前状态采样动作
"""
state_tensor = self._to_tensor(state)
single_state = state_tensor.ndim == 1
if single_state:
state_tensor = state_tensor.unsqueeze(0)
with torch.no_grad():
dist = self.actor.evaluate(state_tensor)
action = dist.sample()
action_np = action.detach().cpu().numpy()
return action_np[0] if single_state else action_np
def get_value(self, state):
state_tensor = self._to_tensor(state)
single_state = state_tensor.ndim == 1
if single_state:
state_tensor = state_tensor.unsqueeze(0)
with torch.no_grad():
values = self.critic(state_tensor).squeeze(-1)
values_np = values.detach().cpu().numpy()
return float(values_np[0]) if single_state else values_np
def _compute_advantages(self, rewards, values, next_values, masks):
"""
使用广义优势估计 (GAE) 计算优势函数
"""
advantages = torch.zeros_like(rewards, device=self.device)
gae = torch.zeros(rewards.size(1), dtype=torch.float32, device=self.device)
for i in reversed(range(rewards.size(0))):
delta = rewards[i] + self.gamma * next_values[i] * masks[i] - values[i]
gae = delta + self.gamma * self.tau * masks[i] * gae
advantages[i] = gae
returns = advantages + values
return returns, advantages
# --- 关键修复 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):
if not memory:
return {}
states = self._to_tensor(np.asarray(memory["states"]))
actions = self._to_tensor(np.asarray(memory["actions"]))
rewards = self._to_tensor(np.asarray(memory["rewards"]))
next_states = self._to_tensor(np.asarray(memory["next_states"]))
masks = self._to_tensor(np.asarray(memory["masks"]))
with torch.no_grad():
rollout_steps, num_envs = rewards.shape
flat_states = states.reshape(rollout_steps * num_envs, -1)
flat_next_states = next_states.reshape(rollout_steps * num_envs, -1)
values = self.critic(flat_states).squeeze(-1).reshape(rollout_steps, num_envs)
next_values = self.critic(flat_next_states).squeeze(-1).reshape(rollout_steps, num_envs)
returns, advantages = self._compute_advantages(rewards, values, next_values, masks)
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
states = flat_states
actions = actions.reshape(rollout_steps * num_envs, -1)
returns = returns.reshape(rollout_steps * num_envs)
advantages = advantages.reshape(rollout_steps * num_envs)
# 【修复】:加大 Critic 的训练力度,从 10 提升到 40 Epochs
# 确保裁判的眼光足够准确
critic_loss_value = 0.0
for _ in range(self.critic_epochs):
critic_loss = F.mse_loss(self.critic(states).squeeze(), returns)
self.critic_optimizer.zero_grad()
critic_loss.backward()
self.critic_optimizer.step()
critic_loss_value += float(critic_loss.detach().item())
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
kl_value = 0.0
surrogate_value = float(surrogate_loss.detach().item())
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
surrogate_value = float(new_surrogate_loss.item())
kl_value = float(kl.item())
break
step_size *= 0.5
kl_value = float(kl.item())
if not success:
vector_to_parameters(old_params, self.actor.parameters())
return {
"critic_loss": critic_loss_value / max(1, self.critic_epochs),
"surrogate_loss": surrogate_value,
"kl": kl_value,
"line_search_success": float(success),
"step_scale": float(step_size if success else 0.0),
}