【发布时间】:2022-01-14 09:06:15
【问题描述】:
我试图在一组参数上运行一个循环,我不想为每个参数创建一个新网络并让它学习几个时期。
目前我的代码如下所示:
def optimize_scale(self, epochs=5, comp_scale=100, scale_list=[1, 100]):
trainer = pyli.Trainer(gpus=1, max_epochs=epochs)
for scale in scale_list:
test_model = CustomNN(num_layers=1, scale=scale, lr=1, pad=True, batch_size=1)
trainer.fit(test_model)
trainer.test(verbose=True)
del test_model
scale_list 的第一个元素一切正常,网络学习了 5 个 epoch 并完成了测试。所有这些都可以在控制台中看到。但是对于scale_list 的所有以下元素,它不起作用,因为旧网络没有被覆盖,而是在调用trainer.fit(model) 时自动加载旧检查点。在控制台中,这通过以下方式指示:
C:\Users\XXXX\AppData\Roaming\Python\Python39\site-packages\pytorch_lightning\callbacks\model_checkpoint.py:623: UserWarning:
Checkpoint directory D:\XXXX\src\lightning_logs\version_0\checkpoints exists and is not empty.
rank_zero_warn(f"Checkpoint directory {dirpath} exists and is not empty.")
train_size = 8 val_size = 1 test_size = 1
Restoring states from the checkpoint path at D:\XXXX\src\lightning_logs\version_0\checkpoints\epoch=4-step=39.ckpt
LOCAL_RANK: 0 - CUDA_VISIBLE_DEVICES: [0]
Loaded model weights from checkpoint at D:\XXXX\src\lightning_logs\version_0\checkpoints\epoch=4-step=39.ckpt
结果是第二个测试输出相同的结果,因为来自旧网络的检查点已加载,它已经完成了所有 5 个 epoch。我虽然添加del test_model 可能有助于完全删除模型,但这不起作用。
在我的搜索中,我发现了一些密切相关的问题,例如:https://github.com/PyTorchLightning/pytorch-lightning/issues/368。但是我没有设法解决我的问题。我认为这与应该覆盖旧网络的新网络具有相同的名称/版本并因此寻找相同的检查点这一事实有关。
如果有人有想法或知道如何规避这一点,我将不胜感激。
【问题讨论】:
标签: python pytorch pytorch-lightning