OOP架构下RL Connect4游戏中torch.save/load使用及训练数据存储咨询
OOP结构下torch.save/load的正确用法与训练数据管理
一、DQNAgent的模型与训练状态保存/加载
在OOP架构中,把保存、加载逻辑封装到DQNAgent类内部是最符合设计原则的做法,既能保证代码内聚性,也方便复用。核心是保存模型的state_dict(而非整个模型对象)、优化器状态,以及训练相关的关键参数(如探索率epsilon、当前训练轮数、累计奖励等),这样既节省空间,又能在不同环境下灵活恢复。
示例代码(封装到DQNAgent类)
import os import torch class DQNAgent: def __init__(self, state_size, action_size): self.state_size = state_size self.action_size = action_size self.epsilon = 1.0 # 初始探索率 self.epsilon_min = 0.01 self.gamma = 0.99 self.current_episode = 0 self.total_rewards = [] # 初始化模型与优化器 self.model = self._build_model() self.optimizer = torch.optim.Adam(self.model.parameters(), lr=0.001) # 自动检测设备(GPU/CPU) self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') self.model.to(self.device) def _build_model(self): # 根据Connect4需求构建DQN模型(示例结构) return torch.nn.Sequential( torch.nn.Linear(self.state_size, 128), torch.nn.ReLU(), torch.nn.Linear(128, 64), torch.nn.ReLU(), torch.nn.Linear(64, self.action_size) ) def save_checkpoint(self, save_path): # 自动创建保存目录 os.makedirs(os.path.dirname(save_path), exist_ok=True) # 打包需要保存的状态 checkpoint = { 'model_state_dict': self.model.state_dict(), 'optimizer_state_dict': self.optimizer.state_dict(), 'epsilon': self.epsilon, 'current_episode': self.current_episode, 'total_rewards': self.total_rewards, 'device': str(self.device) } torch.save(checkpoint, save_path) def load_checkpoint(self, load_path): # 加载时自动适配设备 checkpoint = torch.load(load_path, map_location=self.device) # 恢复模型与优化器状态 self.model.load_state_dict(checkpoint['model_state_dict']) self.optimizer.load_state_dict(checkpoint['optimizer_state_dict']) # 恢复训练参数 self.epsilon = checkpoint['epsilon'] self.current_episode = checkpoint['current_episode'] self.total_rewards = checkpoint['total_rewards'] # 根据场景切换模式:继续训练用train(),测试用eval() self.model.train()
使用方式
- 保存:训练过程中定期保存(比如每100轮)
agent = DQNAgent(state_size=42, action_size=7) # Connect4棋盘为6x7=42维状态 # 训练循环... if agent.current_episode % 100 == 0: agent.save_checkpoint(f'./connect4_checkpoints/agent_episode_{agent.current_episode}.pth')
- 加载:恢复训练或测试预训练模型
agent = DQNAgent(state_size=42, action_size=7) agent.load_checkpoint('./connect4_checkpoints/agent_episode_500.pth') # 若仅需测试,切换为评估模式 agent.model.eval()
二、训练数据的保存与加载位置建议
1. 目录结构规划
建议单独创建一个connect4_checkpoints目录,按类型分类存放文件,方便管理:
connect4_checkpoints/ ├── models/ # 保存Agent的checkpoint文件 ├── logs/ # 保存训练日志(奖励、损失等) └── replay_buffers/ # 按需保存经验回放池数据
2. 训练日志保存
训练过程中的奖励、损失等指标可以用CSV格式保存,方便后续分析:
import csv def save_training_log(log_path, rewards, losses): os.makedirs(os.path.dirname(log_path), exist_ok=True) with open(log_path, 'w', newline='') as f: writer = csv.writer(f) writer.writerow(['Episode', 'Total Reward', 'Average Loss']) for idx in range(len(rewards)): writer.writerow([idx+1, rewards[idx], losses[idx]]) # 使用示例 save_training_log('./connect4_checkpoints/logs/training_log.csv', agent.total_rewards, loss_records)
3. 经验回放池保存
如果需要断点续训时保留回放池数据,可以用pickle序列化保存(注意:若回放池容量过大,会占用较多磁盘空间,建议定期清理或只保存最近的部分数据):
import pickle class ReplayBuffer: def __init__(self, capacity): self.capacity = capacity self.buffer = [] self.position = 0 def push(self, state, action, reward, next_state, done): if len(self.buffer) < self.capacity: self.buffer.append(None) self.buffer[self.position] = (state, action, reward, next_state, done) self.position = (self.position + 1) % self.capacity def save_buffer(self, save_path): os.makedirs(os.path.dirname(save_path), exist_ok=True) with open(save_path, 'wb') as f: pickle.dump(self.buffer, f) def load_buffer(self, load_path): with open(load_path, 'rb') as f: self.buffer = pickle.load(f) self.position = len(self.buffer) % self.capacity
三、关键注意事项
- 优先使用
state_dict:直接保存模型对象可能会因类定义变化导致加载失败,state_dict仅保存参数,兼容性更强。 - 设备适配:加载时通过
map_location自动适配当前设备,避免GPU训练的模型在CPU环境下加载出错。 - 模式切换:加载模型后,若继续训练需调用
model.train(),若用于测试需调用model.eval()(关闭dropout、批量归一化等训练模式特有的层)。
内容的提问来源于stack exchange,提问作者Julie
相关产品推荐
相关产品推荐

