PyTorch模型断点续训:如何从保存的epoch启动训练
PyTorch 断点续训:从保存的epoch恢复训练
你已经实现了write_checkpoint和load_checkpoint函数用于保存/加载模型、优化器、调度器状态及epoch,但当前训练流程无法从保存的epoch位置继续执行,以下是具体修正思路和代码调整方案:
核心问题分析
原代码将模型加载逻辑放在训练循环内部,且循环固定从0开始迭代,同时scheduler加载方式存在错误,导致无法正确延续之前的训练进度。
具体修改步骤
1. 将模型加载逻辑移至训练循环前
把加载操作放在循环开始前完成,避免干扰epoch计数,同时增加路径输入的灵活性:
# 询问是否加载模型 load_model = input('Load a model? (y/n)') start_epoch = 0 if load_model.lower() == 'y': checkpoint_path = input('Enter checkpoint path: ') model, optimizer, start_epoch, scheduler = load_checkpoint(model=model, scheduler=scheduler, optimizer=optimizer, filename=checkpoint_path) # 确保优化器状态迁移到当前设备 for state in optimizer.state.values(): for k, v in state.items(): if isinstance(v, torch.Tensor): state[k] = v.to(device)
2. 修正训练循环的起始范围
将循环从start_epoch开始,而非固定从0启动,保证延续之前的训练进度:
# 从保存的epoch开始训练 for epoch in range(start_epoch, num_epochs): # 训练、评估逻辑...
3. 修复load_checkpoint中的scheduler加载错误
原代码直接替换scheduler对象的方式会导致状态不匹配,应使用load_state_dict方法加载:
def load_checkpoint(model, optimizer, scheduler, filename='/content/checkpoint.pth'): start_epoch = 0 if os.path.isfile(filename): print("=> loading checkpoint '{}'".format(filename)) checkpoint = torch.load(filename) start_epoch = checkpoint['epoch'] model.load_state_dict(checkpoint['state_dict']) optimizer.load_state_dict(checkpoint['optimizer']) # 修正scheduler状态加载方式 scheduler.load_state_dict(checkpoint['scheduler']) print("=> loaded checkpoint '{}' (epoch {})".format(filename, checkpoint['epoch'])) else: print("=> no checkpoint found at '{}'".format(filename)) return model, optimizer, start_epoch, scheduler
4. 调整checkpoint保存时机
原代码在循环内频繁保存并立即加载的逻辑不符合断点续训需求,建议改为每个epoch结束后保存:
# 在epoch训练+评估完成后保存当前状态 write_checkpoint(model=model, epoch=epoch, scheduler=scheduler, optimizer=optimizer)
修改后的完整train_model函数
def train_model(model, train_loader, test_loader, device, learning_rate=1e-1, num_epochs=200): criterion = nn.CrossEntropyLoss() model.to(device) optimizer = optim.SGD(model.parameters(), lr=learning_rate, momentum=0.9, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, milestones=[65, 75], gamma=0.75, last_epoch=-1) # 初始评估 model.eval() eval_loss, eval_accuracy = evaluate_model(model=model, test_loader=test_loader, device=device, criterion=criterion) print("Epoch: {:02d} Eval Loss: {:.3f} Eval Acc: {:.3f}".format(-1, eval_loss, eval_accuracy)) # 加载模型逻辑 load_model = input('Load a model? (y/n)') start_epoch = 0 if load_model.lower() == 'y': checkpoint_path = input('Enter checkpoint path: ') model, optimizer, start_epoch, scheduler = load_checkpoint(model=model, scheduler=scheduler, optimizer=optimizer, filename=checkpoint_path) # 迁移优化器状态到设备 for state in optimizer.state.values(): for k, v in state.items(): if isinstance(v, torch.Tensor): state[k] = v.to(device) # 从保存的epoch开始训练 for epoch in range(start_epoch, num_epochs): model.train() running_loss = 0 running_corrects = 0 for inputs, labels in train_loader: inputs = torch.FloatTensor(inputs).to(device) labels = labels.to(device) optimizer.zero_grad() outputs = model(inputs) _, preds = torch.max(outputs, 1) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() * inputs.size(0) running_corrects += torch.sum(preds == labels.data) train_loss = running_loss / len(train_loader.dataset) train_accuracy = running_corrects / len(train_loader.dataset) # 评估 model.eval() eval_loss, eval_accuracy = evaluate_model(model=model, test_loader=test_loader, device=device, criterion=criterion) scheduler.step() print("Epoch: {:03d} Train Loss: {:.3f} Train Acc: {:.3f} Eval Loss: {:.3f} Eval Acc: {:.3f}".format(epoch, train_loss, train_accuracy, eval_loss, eval_accuracy)) # 保存当前epoch的checkpoint write_checkpoint(model=model, epoch=epoch, scheduler=scheduler, optimizer=optimizer) return model
内容的提问来源于stack exchange,提问作者Manu Dwivedi
相关产品推荐
相关产品推荐

