Files
RL_SAC/train.py
T

87 lines
3.5 KiB
Python
Raw Normal View History

2026-04-04 15:51:08 +08:00
import gymnasium as gym
import numpy as np
import torch
from algorithms.sac import SAC
from utils.replay_buffer import ReplayBuffer
import yaml
import os
# 读取 YAML 配置文件
config_path = os.path.join(os.path.dirname(__file__), 'configs', 'pendulum_config.yaml')
with open(config_path, "r", encoding="utf-8") as f:
config = yaml.safe_load(f)
def main():
# ---------------------------------------------------------
# 2. 实例化环境与获取维度信息
# ---------------------------------------------------------
env = gym.make('Pendulum-v1')
# 动态获取环境的维度信息,确保算法通用性
state_dim = env.observation_space.shape[0]
action_dim = env.action_space.shape[0]
max_action = float(env.action_space.high[0])
print(f"环境已加载: 状态维度 {state_dim}, 动作维度 {action_dim}, 最大动作界限 {max_action}")
# ---------------------------------------------------------
# 3. 实例化 SAC 代理与经验回放池
# ---------------------------------------------------------
agent = SAC(state_dim, action_dim, max_action, config)
replay_buffer = ReplayBuffer(state_dim, action_dim, max_size=config['buffer_size'])
total_steps = 0 # 记录与环境交互的总步数
# ---------------------------------------------------------
# 4. 训练大循环
# ---------------------------------------------------------
for episode in range(config['max_episodes']):
# 重置环境,获取初始状态 (Gymnasium 返回 state 和 info)
state, _ = env.reset()
episode_reward = 0
for step in range(config['max_steps']):
# --- a. 动作选择策略 ---
# 强化学习工程技巧:在训练初期使用纯随机动作,收集高多样性的初始数据
if total_steps < config['start_steps']:
action = env.action_space.sample()
else:
# 预热结束后,交由 SAC 的策略网络进行带噪声的采样
action = agent.select_action(state, evaluate=False)
# --- b. 与环境交互 ---
next_state, reward, terminated, truncated, _ = env.step(action)
# 判断回合是否真正结束 (超时截断 truncated 不算做环境动力学意义上的 done)
done = float(terminated)
# --- c. 存入经验池 ---
replay_buffer.add(state, action, reward, next_state, done)
state = next_state
episode_reward += reward
total_steps += 1
# --- d. 核心学习逻辑 ---
# 只有当经验池里的数据量足够凑齐一个 Batch 时,才开始更新网络
if replay_buffer.size > config['batch_size']:
agent.update(replay_buffer, config['batch_size'])
# 如果提前倒地或撞毁,结束当前回合
if terminated or truncated:
break
# 打印当前回合的训练结果
print(f"Episode: {episode+1:03d} | Total Steps: {total_steps:06d} | Reward: {episode_reward:.2f}")
# --- 阶段性保存模型 (可选) ---
if (episode + 1) % 50 == 0:
torch.save(agent.actor.state_dict(), f"sac_actor_pendulum_ep{episode+1}.pth")
print(f"[*] 已保存第 {episode+1} 回合的模型权重。")
env.close()
print("训练结束!")
if __name__ == "__main__":
main()