Keras中无法从检查点模型恢复训练的问题咨询
Hey there! Let's dig into why your saved model isn't resuming training properly—this is a super common gotcha, and I’ve worked through it plenty of times. Here are the key fixes and checks to get your training back on track:
Most folks start by saving just model.state_dict(), but resuming training needs more than just the model’s weights. You also need to preserve the optimizer’s state (like momentum values, learning rate history) and training metadata (current epoch, best accuracy so far).
Here’s how to save the full checkpoint correctly:
# Inside your training loop, when you hit a better accuracy checkpoint = { 'epoch': epoch, # Current epoch number 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'best_acc': best_acc, # Add this if you're using a learning rate scheduler 'lr_scheduler_state_dict': lr_scheduler.state_dict() } torch.save(checkpoint, 'best_model_checkpoint.pth')
And when loading to resume:
# Load the checkpoint checkpoint = torch.load('best_model_checkpoint.pth') # Restore model, optimizer, and scheduler states model.load_state_dict(checkpoint['model_state_dict']) optimizer.load_state_dict(checkpoint['optimizer_state_dict']) lr_scheduler.load_state_dict(checkpoint['lr_scheduler_state_dict']) # Pick up training from the right epoch start_epoch = checkpoint['epoch'] best_acc = checkpoint['best_acc'] # Critical: Switch model back to training mode (it defaults to eval after load!) model.train()
When you load a model with torch.load(), it automatically switches to evaluation mode (model.eval()). This changes the behavior of layers like BatchNorm and Dropout, which will break your training if you don’t flip it back:
model.train() # Do this right after loading the model state
If you’re using a scheduler (like ReduceLROnPlateau or StepLR), skipping its state means your learning rate will reset to the initial value instead of continuing from where you left off. Always include the scheduler’s state in your checkpoint, like in the first code example.
If you saved the model on a GPU but are loading on CPU (or vice versa), you need to specify the device during loading to avoid errors:
# Load to CPU checkpoint = torch.load('best_model_checkpoint.pth', map_location=torch.device('cpu')) # Load to a specific GPU checkpoint = torch.load('best_model_checkpoint.pth', map_location='cuda:0')
After loading, make sure your optimizer’s tensors are on the correct device too:
model.to(device) for state in optimizer.state.values(): for k, v in state.items(): if isinstance(v, torch.Tensor): state[k] = v.to(device)
Make sure your code is actually saving the checkpoint when it should. For example, if you’re tracking best_acc, confirm you’re updating it correctly every epoch:
current_acc = validate(model, val_loader) if current_acc > best_acc: best_acc = current_acc # Don't forget to update this! torch.save(checkpoint, 'best_model_checkpoint.pth')
If best_acc gets reset every epoch (e.g., it’s defined inside the loop instead of outside), your save logic won’t work as expected.
内容的提问来源于stack exchange,提问作者user41986

