重构 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 torch
import torch.nn as nn
import torch.nn.functional as F
from torch.distributions import Normal
from networks import PolicyNetwork, ValueNetwork
@@ -24,26 +25,26 @@ class PPOAgent:
action_bound,
hidden_dim=128,
gamma=0.99,
tau=0.97,
lr=3e-4,
tau=0.95,
actor_lr=3e-4,
critic_lr=1e-3,
clip_eps=0.2,
k_epochs=10,
minibatch_size=64,
critic_epochs=10,
entropy_coef=0.0,
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.entropy_coef = entropy_coef # 熵正则化系数(可选,鼓励探索)
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=lr)
self.critic_optimizer = torch.optim.Adam(self.critic.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=critic_lr)
def get_action(self, state):
state_tensor = torch.FloatTensor(state).unsqueeze(0)
@@ -52,9 +53,15 @@ class PPOAgent:
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
@@ -66,17 +73,12 @@ class PPOAgent:
def update(self, memory):
"""
PPO-Clip 更新
PPO-Clip 更新 — 对齐论文 Algorithm 1 和 Eq.(9)
Algorithm 1 (Schulman et al. 2017):
for iteration=1, 2, ... do
for actor=1, 2, ..., N do
Run policy π_θold in environment for T timesteps
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. 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]))
@@ -97,24 +99,16 @@ class PPOAgent:
advantages = returns - values
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
# ---------- 3. 训练 Critic (Value Network) ----------
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
# ---------- 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):
# 每轮随机打乱,分成多个 minibatch
np.random.shuffle(indices)
for start in range(0, dataset_size, self.minibatch_size):
@@ -124,27 +118,29 @@ class PPOAgent:
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)
# 概率比 r_t(θ) = π_θ(a|s) / π_θold(a|s)
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
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * mb_advantages
policy_loss = -torch.min(surr1, surr2).mean()
# ---------------------------------
# 可选:熵正则化(鼓励探索)
entropy_loss = -self.entropy_coef * dist.entropy().mean() if self.entropy_coef > 0 else 0
total_loss = policy_loss + entropy_loss
# --- Critic: 价值函数 MSE 损失 ---
value_pred = self.critic(mb_states).squeeze()
critic_loss = F.mse_loss(value_pred, mb_returns)
# Actor 更新(带梯度裁剪,防止策略大幅跳变)
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()
# 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
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.tau = tau
self.kl_margin = kl_margin
@@ -26,6 +26,10 @@ class TRPOAgent:
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) 计算优势函数
@@ -158,82 +162,3 @@ class TRPOAgent:
if not success:
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 numpy as np
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 = []
state, _ = env.reset()
memory = []
current_ep_reward = 0
episodes_completed = 0
print("开始训练 PPO 智能体...")
step_count = 0
while episodes_completed < num_episodes:
action = agent.get_action(state)
# 交互
next_state, reward, terminated, truncated, _ = env.step(action)
# 【极其关键的修复】:只有真正死亡 (terminated) 才清零未来价值
# 绝对不能把时间截断 (truncated) 算作 mask=0
mask = 0.0 if terminated else 1.0
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
current_ep_reward += reward
step_count += 1
@@ -57,15 +47,70 @@ def main():
step_count = 0
env.close()
return episode_rewards
plt.figure(figsize=(10, 5))
plt.plot(episode_rewards)
plt.title('PPO Learning Curve on Pendulum-v1')
plt.xlabel('Episode')
plt.ylabel('Total Reward')
plt.grid(True)
plt.savefig('ppo_learning_curve_final.png')
def smooth(rewards, window=10):
"""滑动平均平滑曲线"""
smoothed = []
for i in range(len(rewards)):
start = max(0, i - window + 1)
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()
print("对比图已保存至 ppo_vs_trpo_comparison.png")
if __name__ == '__main__':
main()
+4 -4
View File
@@ -36,11 +36,11 @@ class PolicyNetwork(nn.Module):
nn.Linear(hidden_dim, hidden_dim),
nn.Tanh(),
nn.Linear(hidden_dim, action_dim),
nn.Tanh() # 关键修复:强制均值输出在 [-1, 1] 之间,防止动作空间爆炸
# 论文原文: tanh 只用于隐藏层激活,输出层是线性的
# 加 tanh 会在 action 接近边界时梯度趋零(饱和),阻碍学习
)
# 初始对数标准差设为 -0.5 (对应的标准差约为 0.6)
# 较小的初始方差有助于防止初期探索步子迈得太大导致系统崩溃
self.action_log_std = nn.Parameter(torch.full((1, action_dim), -0.5))
# 初始对数标准差设为 0(std=1.0),提供充足的初始探索
self.action_log_std = nn.Parameter(torch.zeros(1, action_dim))
def forward(self, state):
# 计算动作均值并将其缩放到实际的物理边界内