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), }