【PyTorch Lightning】.ckpt 是什么?里面有什么?

  1. 什么是检查点(checkpoint, ckpt)?

当模型在训练过程中时,随着其不断接收更多数据,其性能也会发生变化。在训练过程中保存模型的状态是一种最佳实践。这样可以在开发模型的过程中,在每个关键点上获得模型的一个版本,即一个检查点。一旦训练完成,您可以使用在训练过程中找到的性能最佳的检查点。

检查点还使得训练在中断的情况下可以从中断的地方恢复。

PyTorch Lightning 检查点在普通的 PyTorch 中完全可用。

  1. .ckpt 检查点文件里面有什么?

一个 Lightning 检查点包含了模型的整个内部状态的转储。与普通的 PyTorch 不同,Lightning 保存了你在最复杂的分布式训练环境中恢复模型所需的一切。

在 Lightning 检查点中,您会找到:

  • 16 位精度训练的缩放因子(如果使用 16 位精度训练)
  • 当前的 epoch
  • 全局步数
  • LightningModule 的 state_dict
  • 所有优化器的状态
  • 所有学习率调度器的状态
  • 所有回调函数的状态(用于有状态回调函数)
  • 数据模块的状态(用于有状态数据模块)
  • 用于创建模型的超参数(初始参数)
  • 用于创建数据模块的超参数(初始参数)
  • 循环的状态
  1. state_dict 是什么?

nn.Module 的模型权重,具体使用方法如下。

Lightning checkpoints 完全兼容普通的 torch nn.Modules。

python 复制代码
checkpoint = torch.load(CKPT_PATH)
print(checkpoint.keys())

例如,假设像下面这样创建了一个 LightningModule:

python 复制代码
class Encoder(nn.Module):
    ...


class Decoder(nn.Module):
    ...


class Autoencoder(L.LightningModule):
    def __init__(self, encoder, decoder, *args, **kwargs):
        super().__init__()
        self.encoder = encoder
        self.decoder = decoder


autoencoder = Autoencoder(Encoder(), Decoder())

一旦autoencoder训练完成,就可以提取出与 torch nn.Module 相关的权重。

python 复制代码
checkpoint = torch.load(CKPT_PATH)
encoder_weights = {k: v for k, v in checkpoint["state_dict"].items() if k.startswith("encoder.")}
decoder_weights = {k: v for k, v in checkpoint["state_dict"].items() if k.startswith("decoder.")}

官方文档:https://lightning.ai/docs/pytorch/stable/common/checkpointing_basic.html

相关推荐
Ven%23 分钟前
如何让后台运行llamafactory-cli webui 即使关掉了ssh远程连接 也在运行
运维·人工智能·chrome·python·ssh·aigc
Jeo_dmy28 分钟前
(七)人工智能进阶之人脸识别:从刷脸支付到智能安防的奥秘,小白都可以入手的MTCNN+Arcface网络
人工智能·计算机视觉·人脸识别·猪脸识别
睡觉狂魔er2 小时前
自动驾驶控制与规划——Project 5: Lattice Planner
人工智能·机器学习·自动驾驶
xm一点不soso2 小时前
ROS2+OpenCV综合应用--11. AprilTag标签码跟随
人工智能·opencv·计算机视觉
caron43 小时前
Python--正则表达式
python·正则表达式
云卓SKYDROID3 小时前
无人机+Ai应用场景!
人工智能·无人机·科普·高科技·云卓科技
是十一月末3 小时前
机器学习之过采样和下采样调整不均衡样本的逻辑回归模型
人工智能·python·算法·机器学习·逻辑回归
小禾家的3 小时前
.NET AI 开发人员库 --AI Dev Gallery简单示例--问答机器人
人工智能·c#·.net
生信碱移3 小时前
万字长文:机器学习的数学基础(易读)
大数据·人工智能·深度学习·线性代数·算法·数学建模·数据分析
疯狂小料3 小时前
Python3刷算法来呀,贪心系列题单
开发语言·python·算法