You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.15 10:18:08