添加 PPO 算法实现及相关配置,更新训练入口以支持 SAC 和 PPO 模式
This commit is contained in:
@@ -60,3 +60,78 @@ class Logger(object):
|
||||
plt.savefig(save_path, dpi=300, facecolor=fig.get_facecolor(), edgecolor='none')
|
||||
plt.close()
|
||||
print(f"[*] Learning curve saved to: {save_path}")
|
||||
|
||||
@staticmethod
|
||||
def plot_comparison(rewards_dict, window=10, save_dir="."):
|
||||
"""
|
||||
生成三张图:
|
||||
1. comparison_curve.png — 在同 episode 范围内对比(截取到最短算法的长度)
|
||||
2. sac_learning_curve.png — SAC 独立完整曲线
|
||||
3. ppo_learning_curve.png — PPO 独立完整曲线
|
||||
|
||||
Args:
|
||||
rewards_dict (dict): { 'SAC': [r1, r2, ...], 'PPO': [r1, r2, ...] }
|
||||
window (int): 滑动平均窗口大小
|
||||
save_dir (str): 图片保存目录
|
||||
"""
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
|
||||
palette = {
|
||||
'SAC': '#1f77b4', # 蓝色
|
||||
'PPO': '#d62728', # 红色
|
||||
}
|
||||
fallback = ['#2ca02c', '#9467bd', '#8c564b']
|
||||
|
||||
# ========== 图 1: 同 episode 对比(截取到最短长度) ==========
|
||||
min_len = min(len(r) for r in rewards_dict.values())
|
||||
fig, ax = plt.subplots(figsize=(12, 7), facecolor='white')
|
||||
ax.set_facecolor('white')
|
||||
color_idx = 0
|
||||
for algo_name, rewards in rewards_dict.items():
|
||||
color = palette.get(algo_name, fallback[color_idx % len(fallback)])
|
||||
color_idx += 1
|
||||
r = rewards[:min_len]
|
||||
eps = list(range(1, min_len + 1))
|
||||
ax.plot(eps, r, linewidth=1.0, color=color, alpha=0.3)
|
||||
if min_len >= window:
|
||||
avg = np.convolve(r, np.ones(window) / window, mode='valid')
|
||||
ax.plot(range(window, min_len + 1), avg,
|
||||
linewidth=2.5, color=color,
|
||||
label=f'{algo_name} ({window}-ep avg)')
|
||||
ax.set_xlabel('Episodes', fontsize=14, fontweight='bold')
|
||||
ax.set_ylabel('Total Reward', fontsize=14, fontweight='bold')
|
||||
ax.set_title(f'SAC vs PPO — Same Episode Range (1-{min_len})', fontsize=16, fontweight='bold')
|
||||
ax.tick_params(axis='both', which='major', labelsize=12)
|
||||
ax.grid(True, linestyle='--', alpha=0.7)
|
||||
ax.legend(fontsize=13, loc='lower right')
|
||||
plt.tight_layout()
|
||||
p = os.path.join(save_dir, "comparison_curve.png")
|
||||
plt.savefig(p, dpi=300, facecolor=fig.get_facecolor(), edgecolor='none')
|
||||
plt.close()
|
||||
print(f"[*] Comparison curve saved to: {p}")
|
||||
|
||||
# ========== 图 2 & 3: 各算法独立完整曲线 ==========
|
||||
for algo_name, rewards in rewards_dict.items():
|
||||
color = palette.get(algo_name, '#333333')
|
||||
fig, ax = plt.subplots(figsize=(10, 6), facecolor='white')
|
||||
ax.set_facecolor('white')
|
||||
eps = list(range(1, len(rewards) + 1))
|
||||
ax.plot(eps, rewards, linewidth=1.0, color=color, alpha=0.3, label='Episode Reward')
|
||||
if len(rewards) >= window:
|
||||
avg = np.convolve(rewards, np.ones(window) / window, mode='valid')
|
||||
ax.plot(range(window, len(rewards) + 1), avg,
|
||||
linewidth=2.5, color=color,
|
||||
label=f'{window}-Episode Moving Average')
|
||||
ax.set_xlabel('Episodes', fontsize=14, fontweight='bold')
|
||||
ax.set_ylabel('Total Reward', fontsize=14, fontweight='bold')
|
||||
ax.set_title(f'{algo_name} Learning Curve — Pendulum-v1 ({len(rewards)} episodes)',
|
||||
fontsize=16, fontweight='bold')
|
||||
ax.tick_params(axis='both', which='major', labelsize=12)
|
||||
ax.grid(True, linestyle='--', alpha=0.7)
|
||||
ax.legend(fontsize=12, loc='lower right')
|
||||
plt.tight_layout()
|
||||
fname = f"{algo_name.lower()}_learning_curve.png"
|
||||
p = os.path.join(save_dir, fname)
|
||||
plt.savefig(p, dpi=300, facecolor=fig.get_facecolor(), edgecolor='none')
|
||||
plt.close()
|
||||
print(f"[*] {algo_name} learning curve saved to: {p}")
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
class RolloutBuffer:
|
||||
"""
|
||||
On-policy 滚动缓冲区,用于 PPO 算法。
|
||||
|
||||
每次 collect() 收集 T 步数据后,调用 compute_returns_and_advantages()
|
||||
计算 GAE 优势估计,然后通过 get_batches() 将数据切分为 mini-batch
|
||||
供 K 轮 epoch 使用,最后 clear() 清空等待下一轮收集。
|
||||
"""
|
||||
|
||||
def __init__(self, state_dim, action_dim, steps_per_update, gamma, gae_lambda, device):
|
||||
self.steps = steps_per_update
|
||||
self.gamma = gamma
|
||||
self.gae_lambda = gae_lambda
|
||||
self.device = device
|
||||
|
||||
# 预分配存储空间
|
||||
self.states = np.zeros((steps_per_update, state_dim), dtype=np.float32)
|
||||
self.actions_unbounded = np.zeros((steps_per_update, action_dim), dtype=np.float32) # clamp 之前的高斯采样值
|
||||
self.log_probs = np.zeros((steps_per_update, 1), dtype=np.float32)
|
||||
self.rewards = np.zeros((steps_per_update, 1), dtype=np.float32)
|
||||
self.dones = np.zeros((steps_per_update, 1), dtype=np.float32)
|
||||
self.values = np.zeros((steps_per_update, 1), dtype=np.float32)
|
||||
|
||||
# 计算后填充
|
||||
self.returns = np.zeros((steps_per_update, 1), dtype=np.float32)
|
||||
self.advantages = np.zeros((steps_per_update, 1), dtype=np.float32)
|
||||
|
||||
self.ptr = 0
|
||||
self.full = False
|
||||
|
||||
def add(self, state, action_unbounded, log_prob, reward, done, value):
|
||||
"""
|
||||
向缓冲区写入一步数据。
|
||||
action_unbounded: 未裁剪的高斯采样值(形状 [action_dim])
|
||||
log_prob: 该步的 log π(a|s)(标量)
|
||||
value: V(s) 的估计值(标量)
|
||||
"""
|
||||
idx = self.ptr
|
||||
self.states[idx] = state
|
||||
self.actions_unbounded[idx] = action_unbounded
|
||||
self.log_probs[idx] = log_prob
|
||||
self.rewards[idx] = reward
|
||||
self.dones[idx] = done
|
||||
self.values[idx] = value
|
||||
self.ptr += 1
|
||||
if self.ptr >= self.steps:
|
||||
self.full = True
|
||||
|
||||
def compute_returns_and_advantages(self, last_value):
|
||||
"""
|
||||
反向遍历轨迹,计算 GAE 优势估计(论文公式11/12)。
|
||||
只对实际填充的 self.ptr 步数据计算,避免无效数据参与。
|
||||
|
||||
Args:
|
||||
last_value: V(s_{T+1}),下一个状态的价值估计(若 episode 结束则为 0)
|
||||
"""
|
||||
n = self.ptr # 实际有效数据条数
|
||||
last_gae = 0.0
|
||||
for t in reversed(range(n)):
|
||||
if t == n - 1:
|
||||
next_non_terminal = 1.0 - self.dones[t]
|
||||
next_value = last_value
|
||||
else:
|
||||
next_non_terminal = 1.0 - self.dones[t]
|
||||
next_value = self.values[t + 1]
|
||||
|
||||
delta = self.rewards[t] + self.gamma * next_value * next_non_terminal - self.values[t]
|
||||
last_gae = delta + self.gamma * self.gae_lambda * next_non_terminal * last_gae
|
||||
self.advantages[t] = last_gae
|
||||
|
||||
# 回报 G_t = Â_t + V(s_t)
|
||||
self.returns[:n] = self.advantages[:n] + self.values[:n]
|
||||
|
||||
# 优势归一化
|
||||
adv = self.advantages[:n]
|
||||
self.advantages[:n] = (adv - adv.mean()) / (adv.std() + 1e-8)
|
||||
|
||||
def get_batches(self, batch_size):
|
||||
"""
|
||||
将实际填充的数据随机打乱后按 batch_size 切片,生成 mini-batch。
|
||||
|
||||
Yields:
|
||||
(states, actions_unbounded, log_probs, returns, advantages) — 均为 Tensor
|
||||
"""
|
||||
n = self.ptr
|
||||
indices = np.random.permutation(n)
|
||||
for start in range(0, n, batch_size):
|
||||
idx = indices[start: start + batch_size]
|
||||
yield (
|
||||
torch.FloatTensor(self.states[idx]).to(self.device),
|
||||
torch.FloatTensor(self.actions_unbounded[idx]).to(self.device),
|
||||
torch.FloatTensor(self.log_probs[idx]).to(self.device),
|
||||
torch.FloatTensor(self.returns[idx]).to(self.device),
|
||||
torch.FloatTensor(self.advantages[idx]).to(self.device),
|
||||
)
|
||||
|
||||
def clear(self):
|
||||
"""清空缓冲区,为下一轮收集做准备。"""
|
||||
self.ptr = 0
|
||||
self.full = False
|
||||
Reference in New Issue
Block a user