2026-04-04 17:50:02 +08:00
|
|
|
|
"""
|
|
|
|
|
|
训练入口 — 支持 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
|
2026-04-04 15:51:08 +08:00
|
|
|
|
import numpy as np
|
|
|
|
|
|
import torch
|
2026-04-04 17:50:02 +08:00
|
|
|
|
import gymnasium as gym
|
|
|
|
|
|
from typing import cast
|
|
|
|
|
|
from gymnasium.spaces import Box
|
|
|
|
|
|
|
2026-04-04 15:51:08 +08:00
|
|
|
|
from algorithms.sac import SAC
|
2026-04-04 17:50:02 +08:00
|
|
|
|
from algorithms.ppo import PPO
|
2026-04-04 15:51:08 +08:00
|
|
|
|
from utils.replay_buffer import ReplayBuffer
|
2026-04-04 17:50:02 +08:00
|
|
|
|
from utils.logger import Logger
|
2026-04-04 15:51:08 +08:00
|
|
|
|
|
|
|
|
|
|
|
2026-04-04 17:50:02 +08:00
|
|
|
|
# ===========================================================================
|
|
|
|
|
|
# 配置加载
|
|
|
|
|
|
# ===========================================================================
|
2026-04-04 15:51:08 +08:00
|
|
|
|
|
2026-04-04 17:50:02 +08:00
|
|
|
|
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__))
|
2026-04-04 20:10:26 +08:00
|
|
|
|
SAC_CFG_PATH = os.path.join(ROOT, "configs", "sac_pendulum_config.yaml")
|
2026-04-04 17:50:02 +08:00
|
|
|
|
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)
|
2026-04-04 15:51:08 +08:00
|
|
|
|
replay_buffer = ReplayBuffer(state_dim, action_dim, max_size=config['buffer_size'])
|
2026-04-04 17:50:02 +08:00
|
|
|
|
logger = Logger()
|
2026-04-04 15:51:08 +08:00
|
|
|
|
|
2026-04-04 17:50:02 +08:00
|
|
|
|
total_steps = 0
|
2026-04-04 15:51:08 +08:00
|
|
|
|
|
|
|
|
|
|
for episode in range(config['max_episodes']):
|
|
|
|
|
|
state, _ = env.reset()
|
2026-04-04 17:50:02 +08:00
|
|
|
|
episode_reward: float = 0.0
|
|
|
|
|
|
|
|
|
|
|
|
for _ in range(config['max_steps']):
|
|
|
|
|
|
# 动作选择:预热期使用纯随机,之后用策略
|
2026-04-04 15:51:08 +08:00
|
|
|
|
if total_steps < config['start_steps']:
|
|
|
|
|
|
action = env.action_space.sample()
|
|
|
|
|
|
else:
|
|
|
|
|
|
action = agent.select_action(state, evaluate=False)
|
2026-04-04 17:50:02 +08:00
|
|
|
|
|
2026-04-04 15:51:08 +08:00
|
|
|
|
next_state, reward, terminated, truncated, _ = env.step(action)
|
|
|
|
|
|
done = float(terminated)
|
2026-04-04 17:50:02 +08:00
|
|
|
|
|
2026-04-04 15:51:08 +08:00
|
|
|
|
replay_buffer.add(state, action, reward, next_state, done)
|
|
|
|
|
|
state = next_state
|
2026-04-04 17:50:02 +08:00
|
|
|
|
episode_reward += float(reward)
|
2026-04-04 15:51:08 +08:00
|
|
|
|
total_steps += 1
|
2026-04-04 17:50:02 +08:00
|
|
|
|
|
2026-04-04 15:51:08 +08:00
|
|
|
|
if replay_buffer.size > config['batch_size']:
|
|
|
|
|
|
agent.update(replay_buffer, config['batch_size'])
|
2026-04-04 17:50:02 +08:00
|
|
|
|
|
2026-04-04 15:51:08 +08:00
|
|
|
|
if terminated or truncated:
|
|
|
|
|
|
break
|
|
|
|
|
|
|
2026-04-04 17:50:02 +08:00
|
|
|
|
logger.record(episode_reward)
|
|
|
|
|
|
print(f"[SAC] Episode: {episode+1:03d} | Steps: {total_steps:06d} | Reward: {episode_reward:.2f}")
|
|
|
|
|
|
|
2026-04-04 15:51:08 +08:00
|
|
|
|
if (episode + 1) % 50 == 0:
|
2026-04-04 17:50:02 +08:00
|
|
|
|
path = os.path.join(ROOT, f"sac_actor_pendulum_ep{episode+1}.pth")
|
|
|
|
|
|
torch.save(agent.actor.state_dict(), path)
|
|
|
|
|
|
print(f" [*] 模型已保存: {path}")
|
2026-04-04 15:51:08 +08:00
|
|
|
|
|
|
|
|
|
|
env.close()
|
2026-04-04 17:50:02 +08:00
|
|
|
|
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)
|
|
|
|
|
|
|
2026-04-04 15:51:08 +08:00
|
|
|
|
|
|
|
|
|
|
if __name__ == "__main__":
|
|
|
|
|
|
main()
|