如何基于PyTorch OCR模型已存断点续训并忽略警告?
解决PyTorch加载断点时的FutureWarning并完成训练
问题分析
你遇到的是PyTorch torch.load的安全警告:当前默认weights_only=False会使用pickle加载任意对象,存在安全风险,未来版本会默认改为weights_only=True。该警告本身不会中断执行,但需调整代码消除警告并确保训练流程正常完成。
解决方案
1. 修改torch.load调用(推荐)
找到你的load_checkpoint函数中调用torch.load的代码行,添加weights_only=True参数,直接消除警告:
checkpoint_data = torch.load(checkpoint_path, weights_only=True)
由于是你自行保存的断点,仅包含模型和优化器的state_dict,使用该参数完全安全。
2. 若添加参数后报错(断点含自定义对象)
如果你的断点中保存了模型/优化器state_dict之外的自定义对象,暂时无法使用weights_only=True,可以手动抑制该警告:
import warnings def load_checkpoint(model, optimizer, checkpoint_path): # 抑制FutureWarning with warnings.catch_warnings(): warnings.filterwarnings("ignore", category=FutureWarning) checkpoint_data = torch.load(checkpoint_path) # 加载参数逻辑(根据你的断点结构调整) model.load_state_dict(checkpoint_data['model_state_dict']) optimizer.load_state_dict(checkpoint_data['optimizer_state_dict']) start_epoch = checkpoint_data['epoch'] + 1 return start_epoch
3. 规范断点保存(避免后续问题)
修改save_checkpoint函数,仅保存必要的state_dict而非整个对象,确保后续加载时可安全使用weights_only=True:
def save_checkpoint(model, optimizer, epoch, save_path): checkpoint = { 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict() } torch.save(checkpoint, f"{save_path}cp_{epoch+1}.pth")
验证训练流程
修改后重新运行训练代码,警告会被消除,训练循环将正常执行到第24个epoch,并自动保存新的断点文件。
内容的提问来源于stack exchange,提问作者Qusai Triboush
相关产品推荐
相关产品推荐

