You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.13 10:15:34