如何每20轮保存PyTorch模型并从断点继续训练?
PyTorch模型定期保存与断点续训修改方案
针对你的需求,我们在代码中加入定期保存训练检查点和断点加载续训的功能,具体修改如下:
1. 新增命令行参数
首先在参数解析部分添加两个关键参数,适配断点续训和自定义保存间隔:
parser = argparse.ArgumentParser(description="Run DeepGLSTM") parser.add_argument("--dataset",type=str,default='davis', help="Dataset Name (davis,kiba,DTC,Metz,ToxCast,Stitch)") parser.add_argument("--epoch", type = int, default = 1000, help="Number of training epochs. Default is 1000." ) parser.add_argument("--lr", type=float, default = 0.0005, help="learning rate", ) parser.add_argument("--batch_size",type=int, default = 128, help = "Number of drug-tareget per batch. Default is 128 for davis.") parser.add_argument("--save_file",type=str, default=None, help="Base name for saving checkpoint files. E.g. 'davis' generates 'davis_epoch_20.pt'") # 新增:指定断点文件路径,用于续训 parser.add_argument("--resume",type=str, default=None, help="Path to checkpoint file for resuming training. E.g. 'davis_epoch_20.pt'") # 新增:自定义模型保存间隔(默认20epoch) parser.add_argument("--save_interval",type=int, default=20, help="Interval (in epochs) for saving checkpoints. Default is 20.") args = parser.parse_args() print(args) main(args)
2. 核心逻辑修改
在main函数中加入断点加载和定期保存逻辑,确保续训时完全恢复训练状态:
断点加载逻辑
# 初始化训练起始epoch start_epoch = 1 # 加载断点(如果指定) if args.resume: checkpoint = torch.load(args.resume, map_location=device) model.load_state_dict(checkpoint['model_state_dict']) optimizer.load_state_dict(checkpoint['optimizer_state_dict']) start_epoch = checkpoint['epoch'] + 1 # 从下一轮开始训练 best_mse = checkpoint['best_mse'] best_ci = checkpoint['best_ci'] best_epoch = checkpoint['best_epoch'] print(f"Resuming training from epoch {start_epoch}, loaded checkpoint: {args.resume}")
定期保存检查点
在训练循环内添加保存逻辑,每到指定间隔epoch时,保存包含模型权重、优化器状态、当前epoch、最佳指标的完整检查点:
for epoch in range(start_epoch, NUM_EPOCHS + 1): hidden,cell = model.init_hidden(batch_size=TRAIN_BATCH_SIZE) train(model, device, train_loader, optimizer, epoch,hidden,cell) G,P = predicting(model, device, test_loader,hidden,cell) ret = [rmse(G,P),mse(G,P),pearson(G,P),spearman(G,P),ci(G,P),get_rm2(G.reshape(G.shape[0],-1),P.reshape(P.shape[0],-1))] # 定期保存检查点 if args.save_file and epoch % args.save_interval == 0: checkpoint_path = f"{args.save_file}_epoch_{epoch}.pt" torch.save({ 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'best_mse': best_mse, 'best_ci': best_ci, 'best_epoch': best_epoch }, checkpoint_path) print(f"Saved checkpoint to {checkpoint_path}") # 原有最佳模型保存逻辑保留 if ret[1] < best_mse: if args.save_file: best_model_path = f"{args.save_file}_best.model" torch.save(model.state_dict(), best_model_path) with open(result_file_name,'w') as f: f.write('rmse,mse,pearson,spearman,ci,rm2\n') f.write(','.join(map(str,ret))) best_epoch = epoch best_mse = ret[1] best_ci = ret[-2] print('rmse improved at epoch ', best_epoch, '; best_mse,best_ci:', best_mse,best_ci,model_st,dataset) else: print(ret[1],'No improvement since epoch ', best_epoch, '; best_mse,best_ci:', best_mse,best_ci,model_st,dataset)
3. 完整修改后的代码
import argparse import numpy as np import pandas as pd import sys, os from random import shuffle import torch import torch.nn as nn from models.gcn import GCNNet from utils import * # training function at each epoch def train(model, device, train_loader, optimizer, epoch,hidden,cell): print('Training on {} samples...'.format(len(train_loader.dataset))) model.train() for batch_idx, data in enumerate(train_loader): data = data.to(device) optimizer.zero_grad() output = model(data,hidden,cell) loss = loss_fn(output, data.y.view(-1, 1).float().to(device)) loss.backward() optimizer.step() if batch_idx % LOG_INTERVAL == 0: print('Train epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format(epoch, batch_idx * len(data.x), len(train_loader.dataset), 100. * batch_idx / len(train_loader), loss.item())) def predicting(model, device, loader,hidden,cell): model.eval() total_preds = torch.Tensor() total_labels = torch.Tensor() print('Make prediction for {} samples...'.format(len(loader.dataset))) with torch.no_grad(): for data in loader: data = data.to(device) output = model(data,hidden,cell) total_preds = torch.cat((total_preds, output.cpu()), 0) total_labels = torch.cat((total_labels, data.y.view(-1, 1).cpu()), 0) return total_labels.numpy().flatten(),total_preds.numpy().flatten() loss_fn = nn.MSELoss() LOG_INTERVAL = 20 def main(args): dataset = args.dataset modeling = [GCNNet] model_st = modeling[0].__name__ cuda_name = "cuda:0" print('cuda_name:', cuda_name) TRAIN_BATCH_SIZE = args.batch_size TEST_BATCH_SIZE = args.batch_size LR = args.lr NUM_EPOCHS = args.epoch print('Learning rate: ', LR) print('Epochs: ', NUM_EPOCHS) # Main program: iterate over different datasets print('\nrunning on ', model_st + '_' + dataset ) processed_data_file_train = 'data/processed/' + dataset + '_train.pt' processed_data_file_test = 'data/processed/' + dataset + '_test.pt' if ((not os.path.isfile(processed_data_file_train)) or (not os.path.isfile(processed_data_file_test))): print('please run create_data.py to prepare data in pytorch format!') else: train_data = TestbedDataset(root='data', dataset=dataset+'_train') test_data = TestbedDataset(root='data', dataset=dataset+'_test') # make data PyTorch mini-batch processing ready train_loader = DataLoader(train_data, batch_size=TRAIN_BATCH_SIZE, shuffle=True,drop_last=True) test_loader = DataLoader(test_data, batch_size=TEST_BATCH_SIZE, shuffle=False,drop_last=True) # training the model device = torch.device(cuda_name if torch.cuda.is_available() else "cpu") model = modeling[0](k1=1,k2=2,k3=3,embed_dim=128,num_layer=1,device=device).to(device) optimizer = torch.optim.Adam(model.parameters(), lr=LR) best_mse = 1000 best_ci = 0 best_epoch = -1 result_file_name = 'result' + model_st + '_' + dataset + '.csv' # 初始化训练起始epoch start_epoch = 1 # 加载断点(如果指定) if args.resume: checkpoint = torch.load(args.resume, map_location=device) model.load_state_dict(checkpoint['model_state_dict']) optimizer.load_state_dict(checkpoint['optimizer_state_dict']) start_epoch = checkpoint['epoch'] + 1 best_mse = checkpoint['best_mse'] best_ci = checkpoint['best_ci'] best_epoch = checkpoint['best_epoch'] print(f"Resuming training from epoch {start_epoch}, loaded checkpoint: {args.resume}") ##TRAIN _ NUM OF EPOCHES for epoch in range(start_epoch, NUM_EPOCHS + 1): hidden,cell = model.init_hidden(batch_size=TRAIN_BATCH_SIZE) train(model, device, train_loader, optimizer, epoch,hidden,cell) G,P = predicting(model, device, test_loader,hidden,cell) ret = [rmse(G,P),mse(G,P),pearson(G,P),spearman(G,P),ci(G,P),get_rm2(G.reshape(G.shape[0],-1),P.reshape(P.shape[0],-1))] # 定期保存检查点 if args.save_file and epoch % args.save_interval == 0: checkpoint_path = f"{args.save_file}_epoch_{epoch}.pt" torch.save({ 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'best_mse': best_mse, 'best_ci': best_ci, 'best_epoch': best_epoch }, checkpoint_path) print(f"Saved checkpoint to {checkpoint_path}") # 原有最佳模型保存逻辑 if ret[1] < best_mse: if args.save_file: best_model_path = f"{args.save_file}_best.model" torch.save(model.state_dict(), best_model_path) with open(result_file_name,'w') as f: f.write('rmse,mse,pearson,spearman,ci,rm2\n') f.write(','.join(map(str,ret))) best_epoch = epoch best_mse = ret[1] best_ci = ret[-2] print('rmse improved at epoch ', best_epoch, '; best_mse,best_ci:', best_mse,best_ci,model_st,dataset) else: print(ret[1],'No improvement since epoch ', best_epoch, '; best_mse,best_ci:', best_mse,best_ci,model_st,dataset) if __name__ == "__main__": parser = argparse.ArgumentParser(description="Run DeepGLSTM") parser.add_argument("--dataset",type=str,default='davis', help="Dataset Name (davis,kiba,DTC,Metz,ToxCast,Stitch)") parser.add_argument("--epoch", type = int, default = 1000, help="Number of training epochs. Default is 1000." ) parser.add_argument("--lr", type=float, default = 0.0005, help="learning rate", ) parser.add_argument("--batch_size",type=int, default = 128, help = "Number of drug-tareget per batch. Default is 128 for davis.") parser.add_argument("--save_file",type=str, default=None, help="Base name for saving checkpoint files. E.g. 'davis' generates 'davis_epoch_20.pt'") parser.add_argument("--resume",type=str, default=None, help="Path to checkpoint file for resuming training. E.g. 'davis_epoch_20.pt'") parser.add_argument("--save_interval",type=int, default=20, help="Interval (in epochs) for saving checkpoints. Default is 20.") args = parser.parse_args() print(args) main(args)
使用示例
- 首次启动训练(每20epoch保存一次):
python your_script.py --dataset davis --epoch 1000 --save_file davis_model
- 从第20epoch续训到1000epoch:
python your_script.py --dataset davis --epoch 1000 --save_file davis_model --resume davis_model_epoch_20.pt
内容的提问来源于stack exchange,提问作者Hossein
相关产品推荐
相关产品推荐

