import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from networks import PolicyNetwork, ValueNetwork class PPOAgent: """ PPO-Clip (Proximal Policy Optimization, Clipped version) Reference: Schulman et al., "Proximal Policy Optimization Algorithms", 2017. 与 TRPO 的核心区别: - 不再使用共轭梯度 + 线搜索求解约束优化 - 用 clip(ratio, 1-ε, 1+ε) 限制策略更新幅度 - 支持多 epoch 的小批量更新(每次从经验池中采样) """ def __init__( self, state_dim, action_dim, action_bound, hidden_dim=128, gamma=0.99, tau=0.95, actor_lr=3e-4, critic_lr=1e-3, clip_eps=0.2, k_epochs=10, minibatch_size=64, max_grad_norm=0.5, device="cpu", ): self.gamma = gamma self.tau = tau self.clip_eps = clip_eps # PPO-Clip 的裁剪范围 ε self.k_epochs = k_epochs # 每次更新对同一批数据的训练轮数 self.minibatch_size = minibatch_size self.max_grad_norm = max_grad_norm # 梯度裁剪阈值 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.actor_optimizer = torch.optim.Adam(self.actor.parameters(), lr=actor_lr) self.critic_optimizer = torch.optim.Adam(self.critic.parameters(), lr=critic_lr) 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 (Generalized Advantage Estimation) Eq.(11): Â_t = δ_t + (γλ)δ_{t+1} + ... + (γλ)^{T-t+1} δ_{T-1} Eq.(12): δ_t = r_t + γV(s_{t+1}) - V(s_t) """ 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 def update(self, memory): """ PPO-Clip 更新 — 对齐论文 Algorithm 1 和 Eq.(9) 关键改动(相比旧版): 1. Actor 和 Critic 在同一个 minibatch 循环内联合训练 L = L^CLIP - c1 * L^VF (Eq.9,c2=0 不加熵) 2. 梯度裁剪 (max_grad_norm) """ if not memory: return {} # ---------- 1. 准备数据 ---------- 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"])) # ---------- 2. GAE 优势估计 ---------- 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) # ---------- 3. 记录旧策略 π_θold ---------- with torch.no_grad(): old_dist = self.actor.evaluate(states) old_log_probs = old_dist.log_prob(actions).sum(dim=1) # ---------- 4. K epochs × minibatch 联合训练 ---------- dataset_size = states.size(0) indices = np.arange(dataset_size) policy_loss_value = 0.0 critic_loss_value = 0.0 entropy_value = 0.0 minibatch_updates = 0 for _ in range(self.k_epochs): np.random.shuffle(indices) for start in range(0, dataset_size, self.minibatch_size): end = start + self.minibatch_size mb_idx = torch.as_tensor(indices[start:end], device=self.device, dtype=torch.long) mb_states = states[mb_idx] mb_actions = actions[mb_idx] mb_advantages = advantages[mb_idx] mb_returns = returns[mb_idx] mb_old_log_probs = old_log_probs[mb_idx] # --- Actor: PPO-Clip 损失 --- dist = self.actor.evaluate(mb_states) log_probs = dist.log_prob(mb_actions).sum(dim=1) ratio = torch.exp(log_probs - mb_old_log_probs) surr1 = ratio * mb_advantages surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * mb_advantages policy_loss = -torch.min(surr1, surr2).mean() # --- Critic: 价值函数 MSE 损失 --- value_pred = self.critic(mb_states).squeeze() critic_loss = F.mse_loss(value_pred, mb_returns) # Actor 更新(带梯度裁剪,防止策略大幅跳变) self.actor_optimizer.zero_grad() policy_loss.backward() nn.utils.clip_grad_norm_(self.actor.parameters(), self.max_grad_norm) self.actor_optimizer.step() # Critic 更新(不裁剪梯度,Adam 自适应处理大梯度即可) self.critic_optimizer.zero_grad() critic_loss.backward() self.critic_optimizer.step() policy_loss_value += float(policy_loss.detach().item()) critic_loss_value += float(critic_loss.detach().item()) entropy_value += float(dist.entropy().sum(dim=1).mean().detach().item()) minibatch_updates += 1 divisor = max(1, minibatch_updates) return { "policy_loss": policy_loss_value / divisor, "critic_loss": critic_loss_value / divisor, "entropy": entropy_value / divisor, }