重构 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
  对比曲线图(原始曲线 + 滑动平均平滑曲线)
This commit is contained in:
2026-04-02 16:36:55 +08:00
parent 96cd594be9
commit 771eba8607
4 changed files with 121 additions and 155 deletions
+34 -38
View File
@@ -1,5 +1,6 @@
import numpy as np import numpy as np
import torch import torch
import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from torch.distributions import Normal from torch.distributions import Normal
from networks import PolicyNetwork, ValueNetwork from networks import PolicyNetwork, ValueNetwork
@@ -24,26 +25,26 @@ class PPOAgent:
action_bound, action_bound,
hidden_dim=128, hidden_dim=128,
gamma=0.99, gamma=0.99,
tau=0.97, tau=0.95,
lr=3e-4, actor_lr=3e-4,
critic_lr=1e-3,
clip_eps=0.2, clip_eps=0.2,
k_epochs=10, k_epochs=10,
minibatch_size=64, minibatch_size=64,
critic_epochs=10, max_grad_norm=0.5,
entropy_coef=0.0,
): ):
self.gamma = gamma self.gamma = gamma
self.tau = tau self.tau = tau
self.clip_eps = clip_eps # PPO-Clip 的裁剪范围 ε self.clip_eps = clip_eps # PPO-Clip 的裁剪范围 ε
self.k_epochs = k_epochs # 每次更新对同一批数据的训练轮数 self.k_epochs = k_epochs # 每次更新对同一批数据的训练轮数
self.minibatch_size = minibatch_size self.minibatch_size = minibatch_size
self.entropy_coef = entropy_coef # 熵正则化系数(可选,鼓励探索) self.max_grad_norm = max_grad_norm # 梯度裁剪阈值
self.actor = PolicyNetwork(state_dim, action_dim, action_bound, hidden_dim) self.actor = PolicyNetwork(state_dim, action_dim, action_bound, hidden_dim)
self.critic = ValueNetwork(state_dim, hidden_dim) self.critic = ValueNetwork(state_dim, hidden_dim)
self.actor_optimizer = torch.optim.Adam(self.actor.parameters(), lr=lr) self.actor_optimizer = torch.optim.Adam(self.actor.parameters(), lr=actor_lr)
self.critic_optimizer = torch.optim.Adam(self.critic.parameters(), lr=lr) self.critic_optimizer = torch.optim.Adam(self.critic.parameters(), lr=critic_lr)
def get_action(self, state): def get_action(self, state):
state_tensor = torch.FloatTensor(state).unsqueeze(0) state_tensor = torch.FloatTensor(state).unsqueeze(0)
@@ -52,9 +53,15 @@ class PPOAgent:
action = dist.sample() action = dist.sample()
return action.squeeze(0).numpy() 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): def _compute_advantages(self, rewards, values, masks):
""" """
GAE (Generalized Advantage Estimation) 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 = [] returns = []
gae = 0 gae = 0
@@ -66,17 +73,12 @@ class PPOAgent:
def update(self, memory): def update(self, memory):
""" """
PPO-Clip 更新 PPO-Clip 更新 — 对齐论文 Algorithm 1 和 Eq.(9)
Algorithm 1 (Schulman et al. 2017): 关键改动(相比旧版):
for iteration=1, 2, ... do 1. Actor 和 Critic 在同一个 minibatch 循环内联合训练
for actor=1, 2, ..., N do L = L^CLIP - c1 * L^VF (Eq.9c2=0 不加熵)
Run policy π_θold in environment for T timesteps 2. 梯度裁剪 (max_grad_norm)
Compute advantage estimates Aˆ1, ..., AˆT
end for
Optimize surrogate L wrt θ, with K epochs and minibatch size M ≤ NT
θold ← θ
end for
""" """
# ---------- 1. 准备数据 ---------- # ---------- 1. 准备数据 ----------
states = torch.FloatTensor(np.array([m[0] for m in memory])) states = torch.FloatTensor(np.array([m[0] for m in memory]))
@@ -97,24 +99,16 @@ class PPOAgent:
advantages = returns - values advantages = returns - values
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8) advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
# ---------- 3. 训练 Critic (Value Network) ---------- # ---------- 3. 记录旧策略 π_θold ----------
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()
# ---------- 4. 训练 Actor (PPO-Clip 损失) ----------
# 在 no_grad 下记录旧策略的对数概率(对应 Algorithm 1 中的 π_θold
with torch.no_grad(): with torch.no_grad():
old_dist = self.actor.evaluate(states) old_dist = self.actor.evaluate(states)
old_log_probs = old_dist.log_prob(actions).sum(dim=1) old_log_probs = old_dist.log_prob(actions).sum(dim=1)
# ---------- 4. K epochs × minibatch 联合训练 ----------
dataset_size = states.size(0) dataset_size = states.size(0)
indices = np.arange(dataset_size) indices = np.arange(dataset_size)
for _ in range(self.k_epochs): for _ in range(self.k_epochs):
# 每轮随机打乱,分成多个 minibatch
np.random.shuffle(indices) np.random.shuffle(indices)
for start in range(0, dataset_size, self.minibatch_size): for start in range(0, dataset_size, self.minibatch_size):
@@ -124,27 +118,29 @@ class PPOAgent:
mb_states = states[mb_idx] mb_states = states[mb_idx]
mb_actions = actions[mb_idx] mb_actions = actions[mb_idx]
mb_advantages = advantages[mb_idx] mb_advantages = advantages[mb_idx]
mb_returns = returns[mb_idx]
mb_old_log_probs = old_log_probs[mb_idx] mb_old_log_probs = old_log_probs[mb_idx]
# 当前策略的对数概率 # --- Actor: PPO-Clip 损失 ---
dist = self.actor.evaluate(mb_states) dist = self.actor.evaluate(mb_states)
log_probs = dist.log_prob(mb_actions).sum(dim=1) log_probs = dist.log_prob(mb_actions).sum(dim=1)
# 概率比 r_t(θ) = π_θ(a|s) / π_θold(a|s)
ratio = torch.exp(log_probs - mb_old_log_probs) ratio = torch.exp(log_probs - mb_old_log_probs)
# ---------- PPO-Clip 核心 ----------
# L^CLIP(θ) = E[min(r_t(θ) * A_t, clip(r_t(θ), 1-ε, 1+ε) * A_t)]
surr1 = ratio * mb_advantages surr1 = ratio * mb_advantages
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * mb_advantages surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * mb_advantages
policy_loss = -torch.min(surr1, surr2).mean() policy_loss = -torch.min(surr1, surr2).mean()
# ---------------------------------
# 可选:熵正则化(鼓励探索) # --- Critic: 价值函数 MSE 损失 ---
entropy_loss = -self.entropy_coef * dist.entropy().mean() if self.entropy_coef > 0 else 0 value_pred = self.critic(mb_states).squeeze()
critic_loss = F.mse_loss(value_pred, mb_returns)
total_loss = policy_loss + entropy_loss
# Actor 更新(带梯度裁剪,防止策略大幅跳变)
self.actor_optimizer.zero_grad() self.actor_optimizer.zero_grad()
total_loss.backward() policy_loss.backward()
nn.utils.clip_grad_norm_(self.actor.parameters(), self.max_grad_norm)
self.actor_optimizer.step() self.actor_optimizer.step()
# Critic 更新(不裁剪梯度,Adam 自适应处理大梯度即可)
self.critic_optimizer.zero_grad()
critic_loss.backward()
self.critic_optimizer.step()
+5 -80
View File
@@ -6,7 +6,7 @@ from torch.nn.utils import parameters_to_vector, vector_to_parameters
from networks import PolicyNetwork, ValueNetwork from networks import PolicyNetwork, ValueNetwork
class TRPOAgent: 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): 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):
self.gamma = gamma self.gamma = gamma
self.tau = tau self.tau = tau
self.kl_margin = kl_margin self.kl_margin = kl_margin
@@ -26,6 +26,10 @@ class TRPOAgent:
action = dist.sample() action = dist.sample()
return action.squeeze(0).numpy() 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): def _compute_advantages(self, rewards, values, masks):
""" """
使用广义优势估计 (GAE) 计算优势函数 使用广义优势估计 (GAE) 计算优势函数
@@ -158,82 +162,3 @@ class TRPOAgent:
if not success: if not success:
vector_to_parameters(old_params, self.actor.parameters()) 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())
+73 -28
View File
@@ -2,41 +2,31 @@ import gymnasium as gym
import matplotlib.pyplot as plt import matplotlib.pyplot as plt
import numpy as np import numpy as np
from agent.ppo import PPOAgent from agent.ppo import PPOAgent
from agent.trpo import TRPOAgent
def main():
env = gym.make('Pendulum-v1')
state_dim = env.observation_space.shape[0]
action_dim = env.action_space.shape[0]
action_bound = float(env.action_space.high[0])
agent = PPOAgent(state_dim=state_dim, action_dim=action_dim, action_bound=action_bound)
num_episodes = 200
batch_size = 2000
def train_agent(agent, env_name, num_episodes=500, batch_size=2000):
"""通用训练函数,适用于 PPO 和 TRPO"""
env = gym.make(env_name)
episode_rewards = [] episode_rewards = []
state, _ = env.reset() state, _ = env.reset()
memory = [] memory = []
current_ep_reward = 0 current_ep_reward = 0
episodes_completed = 0 episodes_completed = 0
print("开始训练 PPO 智能体...")
step_count = 0 step_count = 0
while episodes_completed < num_episodes: while episodes_completed < num_episodes:
action = agent.get_action(state) action = agent.get_action(state)
# 交互
next_state, reward, terminated, truncated, _ = env.step(action) next_state, reward, terminated, truncated, _ = env.step(action)
# 【极其关键的修复】:只有真正死亡 (terminated) 才清零未来价值
# 绝对不能把时间截断 (truncated) 算作 mask=0
mask = 0.0 if terminated else 1.0
done = terminated or truncated done = terminated or truncated
memory.append([state, action, reward, next_state, mask])
mask = 0.0 if done else 1.0
reward_store = reward
if truncated and not terminated:
reward_store = reward + agent.gamma * agent.get_value(next_state)
memory.append([state, action, reward_store, next_state, mask])
state = next_state state = next_state
current_ep_reward += reward current_ep_reward += reward
step_count += 1 step_count += 1
@@ -57,15 +47,70 @@ def main():
step_count = 0 step_count = 0
env.close() env.close()
return episode_rewards
plt.figure(figsize=(10, 5))
plt.plot(episode_rewards) def smooth(rewards, window=10):
plt.title('PPO Learning Curve on Pendulum-v1') """滑动平均平滑曲线"""
plt.xlabel('Episode') smoothed = []
plt.ylabel('Total Reward') for i in range(len(rewards)):
plt.grid(True) start = max(0, i - window + 1)
plt.savefig('ppo_learning_curve_final.png') smoothed.append(np.mean(rewards[start:i + 1]))
return smoothed
def main():
env_name = 'Pendulum-v1'
env = gym.make(env_name)
state_dim = env.observation_space.shape[0]
action_dim = env.action_space.shape[0]
action_bound = float(env.action_space.high[0])
env.close()
num_episodes = 500
# --- 训练 PPO ---
print("=" * 50)
print("开始训练 PPO 智能体...")
print("=" * 50)
ppo_agent = PPOAgent(state_dim=state_dim, action_dim=action_dim, action_bound=action_bound)
ppo_rewards = train_agent(ppo_agent, env_name, num_episodes)
# --- 训练 TRPO ---
print("=" * 50)
print("开始训练 TRPO 智能体...")
print("=" * 50)
trpo_agent = TRPOAgent(state_dim=state_dim, action_dim=action_dim, action_bound=action_bound)
trpo_rewards = train_agent(trpo_agent, env_name, num_episodes)
# --- 对比画图 ---
fig, axes = plt.subplots(1, 2, figsize=(16, 6))
# 左图:原始奖励曲线
axes[0].plot(ppo_rewards, alpha=0.3, color='blue', label='PPO (raw)')
axes[0].plot(trpo_rewards, alpha=0.3, color='red', label='TRPO (raw)')
axes[0].plot(smooth(ppo_rewards, 20), color='blue', linewidth=2, label='PPO (smooth)')
axes[0].plot(smooth(trpo_rewards, 20), color='red', linewidth=2, label='TRPO (smooth)')
axes[0].set_title('PPO vs TRPO on Pendulum-v1')
axes[0].set_xlabel('Episode')
axes[0].set_ylabel('Total Reward')
axes[0].legend()
axes[0].grid(True)
# 右图:滑动平均对比(更清晰)
axes[1].plot(smooth(ppo_rewards, 20), color='blue', linewidth=2, label='PPO')
axes[1].plot(smooth(trpo_rewards, 20), color='red', linewidth=2, label='TRPO')
axes[1].set_title('PPO vs TRPO (Smoothed, window=20)')
axes[1].set_xlabel('Episode')
axes[1].set_ylabel('Total Reward')
axes[1].legend()
axes[1].grid(True)
plt.tight_layout()
plt.savefig('ppo_vs_trpo_comparison.png', dpi=150)
plt.show() plt.show()
print("对比图已保存至 ppo_vs_trpo_comparison.png")
if __name__ == '__main__': if __name__ == '__main__':
main() main()
+4 -4
View File
@@ -36,11 +36,11 @@ class PolicyNetwork(nn.Module):
nn.Linear(hidden_dim, hidden_dim), nn.Linear(hidden_dim, hidden_dim),
nn.Tanh(), nn.Tanh(),
nn.Linear(hidden_dim, action_dim), nn.Linear(hidden_dim, action_dim),
nn.Tanh() # 关键修复:强制均值输出在 [-1, 1] 之间,防止动作空间爆炸 # 论文原文: tanh 只用于隐藏层激活,输出层是线性的
# 加 tanh 会在 action 接近边界时梯度趋零(饱和),阻碍学习
) )
# 初始对数标准差设为 -0.5 (对应的标准差约为 0.6) # 初始对数标准差设为 0(std=1.0),提供充足的初始探索
# 较小的初始方差有助于防止初期探索步子迈得太大导致系统崩溃 self.action_log_std = nn.Parameter(torch.zeros(1, action_dim))
self.action_log_std = nn.Parameter(torch.full((1, action_dim), -0.5))
def forward(self, state): def forward(self, state):
# 计算动作均值并将其缩放到实际的物理边界内 # 计算动作均值并将其缩放到实际的物理边界内