添加 PPO 算法实现及相关配置,更新训练入口以支持 SAC 和 PPO 模式

This commit is contained in:
2026-04-04 17:50:02 +08:00
parent f3d2d1a85f
commit a2ce5073c5
10 changed files with 675 additions and 73 deletions
+75
View File
@@ -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}")