在Google Colab中基于仅含权重的ckpt恢复PyTorch Lightning训练
解决PyTorch Lightning从仅权重Checkpoint恢复训练的问题
问题描述
在Google Colab中尝试从仅保存权重的checkpoint文件恢复PyTorch Lightning训练时,直接使用trainer.fit(ckpt_path=...)报错,提示文件仅包含权重。尝试加载该ckpt后执行torch.save(model,"model.ckpt")保存,仍然无效,报错信息如下:
rank_zero_warn(f"Checkpoint directory {dirpath} exists and is not empty.") 2022-11-12 02:03:37,822 - Restoring states from the checkpoint path at /content/drive/MyDrive/maicon_qualifiers/bit_fold3_cosine_epoch=00048_val_mIoU=0.5898_val_loss=0.0447.ckpt Traceback (most recent call last): File "train.py", line 38, in <module> trainer.fit(model, dm, ckpt_path=opt.resume_from_checkpoint) File "/usr/local/lib/python3.7/dist-packages/pytorch_lightning/trainer/trainer.py", line 583, in fit self, self._fit_impl, model, train_dataloaders, val_dataloaders, datamodule, ckpt_path File "/usr/local/lib/python3.7/dist-packages/pytorch_lightning/trainer/call.py", line 38, in _call_and_handle_interrupt return trainer_fn(*args, **kwargs) File "/usr/local/lib/python3.7/dist-packages/pytorch_lightning/trainer/trainer.py", line 624, in _fit_impl self._run(model, ckpt_path=self.ckpt_path) File "/usr/local/lib/python3.7/dist-packages/pytorch_lightning/trainer/trainer.py", line 1005, in _run self._restore_modules_and_callbacks(ckpt_path) File "/usr/local/lib/python3.7/dist-packages/pytorch_lightning/trainer/trainer.py", line 959, in _restore_modules_and_callbacks self._checkpoint_connector.resume_start(checkpoint_path) File "/usr/local/lib/python3.7/dist-packages/pytorch_lightning/trainer/connectors/checkpoint_connector.py", line 89, in resume_start self._loaded_checkpoint = self._load_and_validate_checkpoint(checkpoint_path) File "/usr/local/lib/python3.7/dist-packages/pytorch_lightning/trainer/connectors/checkpoint_connector.py", line 94, in _load_and_validate_checkpoint if any(key in loaded_checkpoint for key in DEPRECATED_CHECKPOINT_KEYS): File "/usr/local/lib/python3.7/dist-packages/pytorch_lightning/trainer/connectors/checkpoint_connector.py", line 94, in <genexpr> if any(key in loaded_checkpoint for key in DEPRECATED_CHECKPOINT_KEYS): TypeError: argument of type 'CDBaseModel' is not iterable
解决方案
1. 转换仅权重文件为PyTorch Lightning兼容的Checkpoint格式
PyTorch Lightning需要的checkpoint是包含state_dict、hyper_parameters等键的字典结构,而非直接保存的模型对象。执行以下步骤转换:
import torch from your_model_module import CDBaseModel # 替换为你的模型类所在模块 # 初始化与训练时结构完全一致的模型 model = CDBaseModel(your_model_hparams) # 传入训练时使用的超参数 # 加载仅权重的checkpoint weight_checkpoint = torch.load("/content/drive/MyDrive/maicon_qualifiers/bit_fold3_cosine_epoch=00048_val_mIoU=0.5898_val_loss=0.0447.ckpt") model.load_state_dict(weight_checkpoint) # 构造Lightning标准checkpoint字典 lightning_ckpt = { "state_dict": model.state_dict(), "hyper_parameters": model.hparams, # 若模型使用hparams则保留,否则可省略 # 若需要恢复训练时的优化器状态,需额外保存(如果之前有记录) # "optimizer_states": [optimizer.state_dict() for optimizer in trainer.optimizers], } # 保存为Lightning可识别的checkpoint torch.save(lightning_ckpt, "full_lightning_ckpt.ckpt")
2. 使用转换后的Checkpoint恢复训练
修改trainer.fit调用,传入新生成的checkpoint路径:
trainer.fit(model, dm, ckpt_path="full_lightning_ckpt.ckpt")
问题原因
直接执行torch.save(model, "model.ckpt")保存的是完整模型对象,而非Lightning期望的字典结构。当Lightning尝试加载该文件时,会将模型对象当作字典去检查键值,从而触发TypeError: argument of type 'CDBaseModel' is not iterable错误。
内容的提问来源于stack exchange,提问作者shersher
相关产品推荐
相关产品推荐

