# ============================================================================== # Author: Hongru Liu # Affiliation: School of Power and Energy, Northwestern Polytechnical University # Version: 1.0 # Contact: hongruliu@mail.nwpu.edu.cn # ============================================================================== import matplotlib.pyplot as plt import numpy as np import os class Logger(object): def __init__(self): """ 初始化日志记录器,用于暂存训练过程中的各项指标 """ self.episode_rewards = [] def record(self, reward): """ 记录每个 Episode 的总奖励 """ self.episode_rewards.append(reward) def plot_learning_curve(self, save_dir="."): """ 绘制并保存学习曲线 (Learning Curve) """ # 确保保存目录存在 os.makedirs(save_dir, exist_ok=True) # 创建画布,严格设置白色背景 fig, ax = plt.subplots(figsize=(10, 6), facecolor='white') ax.set_facecolor('white') # 绘制奖励曲线,使用加粗线条以满足论文发表的视觉要求 ax.plot(self.episode_rewards, linewidth=2.5, color='#1f77b4', label='Episode Reward') # 计算并绘制 10 个 Episode 的滑动平均线,让趋势更清晰 if len(self.episode_rewards) >= 10: moving_avg = np.convolve(self.episode_rewards, np.ones(10)/10, mode='valid') ax.plot(range(9, len(self.episode_rewards)), moving_avg, linewidth=2.5, color='#ff7f0e', label='10-Episode Moving Average') # 设置全英文的坐标轴标签和图例,调整字体大小 ax.set_xlabel('Episodes', fontsize=14, fontweight='bold') ax.set_ylabel('Total Reward', fontsize=14, fontweight='bold') ax.set_title('Training Learning Curve (Pendulum-v1)', 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() save_path = os.path.join(save_dir, "learning_curve.png") 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}")