PyTorch Lightning加载checkpoint报IsADirectoryError排查
问题描述
请问如下训练函数为何会触发IsADirectoryError报错:
def train_graph_classifier(model_name, **model_kwargs): pl.seed_everything(42) # 创建带生成回调的PyTorch Lightning训练器 root_dir = os.path.join('/home/predictor2', "GraphLevel" + model_name) os.makedirs(root_dir, exist_ok=True) trainer = pl.Trainer(default_root_dir=root_dir, callbacks=[ModelCheckpoint(save_weights_only=True, mode="max", monitor="val_acc")], gpus=1 if str(device).startswith("cuda") else 0, max_epochs=500, progress_bar_refresh_rate=0) trainer.logger._default_hp_metric = None # 不需要的可选日志参数 # 检查是否存在预训练模型,存在则直接加载跳过训练 pretrained_filename = os.path.join('/home/predictor2', f"GraphLevel{model_name}.ckpt") if os.path.isfile(pretrained_filename): print("Found pretrained model, loading...") model = GraphLevelGNN.load_from_checkpoint(pretrained_filename) else: pl.seed_everything(42) model = GraphLevelGNN(c_in=dataset.num_node_features, c_out=1 if dataset.num_classes==2 else dataset.num_classes, #change **model_kwargs) trainer.fit(model, graph_train_loader, graph_val_loader) model = GraphLevelGNN.load_from_checkpoint(trainer.checkpoint_callback.best_model_path) # 在验证集、测试集上测试最优模型 train_result = trainer.test(model, graph_train_loader, verbose=False) test_result = trainer.test(model, graph_test_loader, verbose=False) result = {"test": test_result[0]['test_acc'], "train": train_result[0]['test_acc']} return model, result
函数运行返回的错误栈如下:
Traceback (most recent call last): File "stability_v3_alternative_net.py", line 604, in <module> dp_rate=0.2) File "stability_v3_alternative_net.py", line 591, in train_graph_classifier model = GraphLevelGNN.load_from_checkpoint(trainer.checkpoint_callback.best_model_path) File "/root/miniconda3/lib/python3.7/site-packages/pytorch_lightning/core/saving.py", line 139, in load_from_checkpoint checkpoint = pl_load(checkpoint_path, map_location=lambda storage, loc: storage) File "/root/miniconda3/lib/python3.7/site-packages/pytorch_lightning/utilities/cloud_io.py", line 46, in load with fs.open(path_or_url, "rb") as f: File "/root/miniconda3/lib/python3.7/site-packages/fsspec/spec.py", line 1043, in open **kwargs, File "/root/miniconda3/lib/python3.7/site-packages/fsspec/implementations/local.py", line 159, in _open return LocalFileOpener(path, mode, fs=self, **kwargs) File "/root/miniconda3/lib/python3.7/site-packages/fsspec/implementations/local.py", line 254, in __init__ self._open() File "/root/miniconda3/lib/python3.7/site-packages/fsspec/implementations/local.py", line 259, in _open self.f = open(self.path, mode=self.mode) IsADirectoryError: [Errno 21] Is a directory: '/home/predictor'
其中/home/predictor是当前工作目录,已特意创建predictor2目录存储训练相关文件,即便把代码中predictor2替换为predictor仍会触发相同错误。已知错误含义是程序读取文件时传入的路径为目录而非目标文件,但代码中没有显式引用工作目录/home/predictor,无法定位触发点,代码参考公开PyTorch GNN教程示例编写,需要排查问题原因与解决方案。
问题根因
报错核心是trainer.checkpoint_callback.best_model_path返回了目录路径/home/predictor而非具体的ckpt文件路径,由两类常见问题触发:
- 检查点保存逻辑未生效:配置的
ModelCheckpoint监控指标为val_acc,如果验证步骤未正常运行、或val_acc指标没有被正确记录到日志中,检查点回调不会生成任何有效模型文件,此时best_model_path会被默认赋值为当前工作目录,加载时自然触发目录错误。 - PyTorch Lightning版本兼容问题:你当前使用的是适配Python3.7的1.x旧版本PL,如果不给
ModelCheckpoint显式指定dirpath参数,回调不会自动继承Trainer的default_root_dir作为检查点存储路径,会回退到当前工作目录,未生成有效检查点时就会直接返回目录路径。
解决方案
按优先级依次排查修改:
- 给
ModelCheckpoint显式指定检查点存储目录,避免路径回退:# 先定义检查点专属存储目录 ckpt_dir = os.path.join(root_dir, "checkpoints") os.makedirs(ckpt_dir, exist_ok=True) # 初始化回调时显式传入路径、文件名规则参数 checkpoint_callback = ModelCheckpoint( dirpath=ckpt_dir, filename="best-model-{epoch:02d}-{val_acc:.3f}", save_weights_only=True, mode="max", monitor="val_acc" ) trainer = pl.Trainer( default_root_dir=root_dir, callbacks=[checkpoint_callback], gpus=1 if str(device).startswith("cuda") else 0, max_epochs=500, progress_bar_refresh_rate=0 ) - 确认验证流程正常,
val_acc指标被正确记录:检查验证步validation_step中是否正确调用self.log("val_acc", acc, prog_bar=True),确保日志记录的指标名和ModelCheckpoint的monitor参数完全一致,无拼写错误。 - 加载最优模型前增加路径校验,提前拦截无效路径问题:
trainer.fit(model, graph_train_loader, graph_val_loader) best_ckpt_path = trainer.checkpoint_callback.best_model_path # 确认路径是有效文件再执行加载 assert os.path.isfile(best_ckpt_path), f"最佳模型路径无效,当前值为{best_ckpt_path},请检查检查点保存逻辑" model = GraphLevelGNN.load_from_checkpoint(best_ckpt_path)
内容的提问来源于stack exchange,提问作者Slowat_Kela
相关产品推荐
相关产品推荐

