Files
RL_TRPO/agent/ppo.py
T
Hongru 771eba8607 重构 PPO/TRPO 训练流程并添加对比绘图
- PPO: 改为 Actor/Critic 联合小批量训练,新增梯度裁剪 (max_grad_norm),
  分离 actor_lr/critic_lr,添加 get_value(),GAE 部分补充论文公式注释
- TRPO: 添加 get_value(),调整 tau 从 0.97 到 0.95
- Networks: 移除 PolicyNet 输出层的 tanh,初始化 log_std=0 以增强探索
- Main: 抽取 train_agent() 通用训练函数,新增 TRPO 训练和 PPO vs TRPO
  对比曲线图(原始曲线 + 滑动平均平滑曲线)
2026-04-02 16:36:55 +08:00

147 lines
5.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.distributions import Normal
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,
):
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.actor = PolicyNetwork(state_dim, action_dim, action_bound, hidden_dim)
self.critic = ValueNetwork(state_dim, hidden_dim)
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 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 get_value(self, state):
with torch.no_grad():
return self.critic(torch.FloatTensor(state).unsqueeze(0)).squeeze().item()
def _compute_advantages(self, rewards, 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)
"""
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
def update(self, memory):
"""
PPO-Clip 更新 — 对齐论文 Algorithm 1 和 Eq.(9)
关键改动(相比旧版):
1. Actor 和 Critic 在同一个 minibatch 循环内联合训练
L = L^CLIP - c1 * L^VF (Eq.9c2=0 不加熵)
2. 梯度裁剪 (max_grad_norm)
"""
# ---------- 1. 准备数据 ----------
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]
# ---------- 2. GAE 优势估计 ----------
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)
# ---------- 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)
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 = indices[start:end]
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()