""" 训练入口 — 支持 SAC / PPO / 对比(compare) 三种模式 用法: python train.py --algo sac # 仅训练 SAC(行为与原始版本一致) python train.py --algo ppo # 仅训练 PPO python train.py --algo compare # 依次训练 SAC 和 PPO,结束后输出对比曲线 可选参数: --env Gymnasium 环境 ID(默认 Pendulum-v1) --seed 随机种子(默认 0) """ import argparse import os import yaml import numpy as np import torch import gymnasium as gym from typing import cast from gymnasium.spaces import Box from algorithms.sac import SAC from algorithms.ppo import PPO from utils.replay_buffer import ReplayBuffer from utils.logger import Logger # =========================================================================== # 配置加载 # =========================================================================== def load_config(path: str) -> dict: with open(path, "r", encoding="utf-8") as f: return yaml.safe_load(f) ROOT = os.path.dirname(os.path.abspath(__file__)) SAC_CFG_PATH = os.path.join(ROOT, "configs", "sac_pendulum_config.yaml") PPO_CFG_PATH = os.path.join(ROOT, "configs", "ppo_pendulum_config.yaml") # =========================================================================== # 工具函数 # =========================================================================== def set_seed(seed: int, env: gym.Env): np.random.seed(seed) torch.manual_seed(seed) env.reset(seed=seed) def make_env(env_id: str): env = gym.make(env_id) obs_space = cast(Box, env.observation_space) act_space = cast(Box, env.action_space) if obs_space.shape is None or act_space.shape is None: raise ValueError("环境的 observation/action space 不支持 shape 维度读取。") state_dim = obs_space.shape[0] action_dim = act_space.shape[0] max_action = float(act_space.high[0]) return env, state_dim, action_dim, max_action # =========================================================================== # SAC 训练循环 (off-policy) # =========================================================================== def train_sac(env_id: str, seed: int) -> list: """ 训练 SAC 并返回每个 episode 的总奖励列表。 """ config = load_config(SAC_CFG_PATH) env, state_dim, action_dim, max_action = make_env(env_id) set_seed(seed, env) print(f"\n{'='*60}") print(f" Training SAC on {env_id}") print(f" state_dim={state_dim}, action_dim={action_dim}, max_action={max_action}") print(f"{'='*60}\n") agent = SAC(state_dim, action_dim, max_action, config) replay_buffer = ReplayBuffer(state_dim, action_dim, max_size=config['buffer_size']) logger = Logger() total_steps = 0 for episode in range(config['max_episodes']): state, _ = env.reset() episode_reward: float = 0.0 for _ in range(config['max_steps']): # 动作选择:预热期使用纯随机,之后用策略 if total_steps < config['start_steps']: action = env.action_space.sample() else: action = agent.select_action(state, evaluate=False) next_state, reward, terminated, truncated, _ = env.step(action) done = float(terminated) replay_buffer.add(state, action, reward, next_state, done) state = next_state episode_reward += float(reward) total_steps += 1 if replay_buffer.size > config['batch_size']: agent.update(replay_buffer, config['batch_size']) if terminated or truncated: break logger.record(episode_reward) print(f"[SAC] Episode: {episode+1:03d} | Steps: {total_steps:06d} | Reward: {episode_reward:.2f}") if (episode + 1) % 50 == 0: path = os.path.join(ROOT, f"sac_actor_pendulum_ep{episode+1}.pth") torch.save(agent.actor.state_dict(), path) print(f" [*] 模型已保存: {path}") env.close() logger.plot_learning_curve(save_dir=ROOT) print("\n[SAC] 训练完成!\n") return logger.episode_rewards # =========================================================================== # PPO 训练循环 (on-policy) # =========================================================================== def train_ppo(env_id: str, seed: int) -> list: """ 训练 PPO 并返回每个 episode 的总奖励列表。 PPO 是 on-policy 的:先收集固定 T 步数据(steps_per_update), 然后用这批数据做 K 轮 epoch 更新,再继续收集。 """ config = load_config(PPO_CFG_PATH) env, state_dim, action_dim, max_action = make_env(env_id) set_seed(seed, env) print(f"\n{'='*60}") print(f" Training PPO on {env_id}") print(f" state_dim={state_dim}, action_dim={action_dim}, max_action={max_action}") print(f"{'='*60}\n") agent = PPO(state_dim, action_dim, max_action, config) logger = Logger() steps_per_update = config.get('steps_per_update', 2048) max_episodes = config['max_episodes'] max_steps = config['max_steps'] total_steps = 0 buffer_steps = 0 episode = 0 episode_reward: float = 0.0 state, _ = env.reset() while episode < max_episodes: action, action_raw, log_prob, value = agent.select_action(state) next_state, reward, terminated, truncated, _ = env.step(action) done = float(terminated) agent.rollout.add(state, action_raw, log_prob, reward, done, value) state = next_state episode_reward += float(reward) total_steps += 1 buffer_steps += 1 # ---- episode 结束 ---- if terminated or truncated: logger.record(episode_reward) print(f"[PPO] Episode: {episode+1:03d} | Steps: {total_steps:06d} | Reward: {episode_reward:.2f}") # 阶段性保存(避免重复保存) if (episode + 1) % 50 == 0: path = os.path.join(ROOT, f"ppo_actor_pendulum_ep{episode+1}.pth") torch.save(agent.actor.state_dict(), path) print(f" [*] 模型已保存: {path}") episode += 1 episode_reward = 0.0 state, _ = env.reset() if episode >= max_episodes: break # ---- 收集够 T 步 → 触发 PPO 更新 ---- if buffer_steps >= steps_per_update: last_value = 0.0 if done else agent.get_value(state) agent.rollout.compute_returns_and_advantages(last_value) info = agent.update() buffer_steps = 0 print(f" [PPO update] actor_loss={info['actor_loss']:.4f} | " f"value_loss={info['value_loss']:.4f} | entropy={info['entropy']:.4f}") env.close() # 若缓冲区中还有剩余数据,做最后一次更新 if buffer_steps > 0: agent.rollout.compute_returns_and_advantages(0.0) agent.update() logger.plot_learning_curve(save_dir=ROOT) print("\n[PPO] 训练完成!\n") return logger.episode_rewards # =========================================================================== # 对比模式:依次训练两者,输出对比曲线 # =========================================================================== def train_compare(env_id: str, seed: int): print("\n" + "="*60) print(" Compare Mode: SAC vs PPO") print("="*60) rewards_sac = train_sac(env_id, seed) rewards_ppo = train_ppo(env_id, seed) Logger.plot_comparison( rewards_dict={'SAC': rewards_sac, 'PPO': rewards_ppo}, window=10, save_dir=ROOT, ) print("\n[Compare] 对比训练完成!已生成 3 张图:") print(" - comparison_curve.png (同 episode 范围对比)") print(" - sac_learning_curve.png (SAC 独立完整曲线)") print(" - ppo_learning_curve.png (PPO 独立完整曲线)") # =========================================================================== # 命令行入口 # =========================================================================== def parse_args(): parser = argparse.ArgumentParser( description="SAC / PPO 强化学习训练脚本 (Pendulum-v1)" ) parser.add_argument( "--algo", type=str, default="sac", choices=["sac", "ppo", "compare"], help="选择训练的算法: sac | ppo | compare (默认: sac)" ) parser.add_argument( "--env", type=str, default="Pendulum-v1", help="Gymnasium 环境 ID (默认: Pendulum-v1)" ) parser.add_argument( "--seed", type=int, default=0, help="随机种子 (默认: 0)" ) return parser.parse_args() def main(): args = parse_args() if args.algo == "sac": train_sac(args.env, args.seed) elif args.algo == "ppo": train_ppo(args.env, args.seed) elif args.algo == "compare": train_compare(args.env, args.seed) if __name__ == "__main__": main()