PyTorch Lightning trainer.fit卡在epoch 0问题排查
PyTorch Lightning 启动训练后卡在Epoch 0无报错问题排查
问题复现场景
从TensorFlow迁移三输入多分支3分类模型到PyTorch + PyTorch Lightning框架时,
trainer.fit()执行后无任何报错输出,控制台可正常打印环境检测信息:识别到CUDA设备0,模型共239K可训练参数、0不可训练参数,总大小约0.958MB,后续训练进度条始终停留在Epoch 0: 0%| | 0/782 [00:00<?, ?it/s],无任何进展。
现有实现说明
- 训练配置:指定任务类型为RC、使用GI4E数据集,配置batch size为16、学习率0.001、训练轮次500、训练集拆分比例0.8,配置项同时用于不同数据集的预处理逻辑分支选择
- 基础数据集类
RCDataset:按配置读取GI4E/BIOID数据集下的图像路径,分别构建非眼部、左眼、右眼三路图像的路径列表,生成对应0/1/2分类标签;实现__getitem__方法完成图像读取、Tensor格式转换,实现__len__方法返回数据集总样本量 - Lightning数据模块
RCDataModule:继承pl.LightningDataModule,按配置比例拆分训练集与验证集,分别生成训练、验证、预测阶段的DataLoader,设置num_workers=12 - 基础模型类
RCBase:搭建三路结构完全一致的卷积神经网络分支,单分支结构包含Conv2d、ReLU、MaxPool2d、Flatten、Linear层,最终经Softmax输出3分类结果;forward方法接收三路图像输入,返回三路分支的预测结果 - Lightning模型封装类
RCPL:继承pl.LightningModule,加载RCBase基础模型,配置Adam优化器;训练步骤计算三路分支交叉熵损失的均值并写入日志,验证步骤计算验证集损失并写入日志,预测步骤返回模型推理结果
排查方向与修复方案
按出现概率从高到低排序:
- DataLoader多进程阻塞(最高发)
当前设置num_workers=12是最常见的卡0进度诱因:Windows系统下DataLoader多进程spawn模式容易出现死锁;如果图像读取用的OpenCV、PIL库在多进程fork时没有正确重新初始化,也会出现无报错阻塞。
修复方式:先把num_workers改成0走单进程加载测试,如果能正常跑通,再根据系统环境逐步调大worker数值;Linux环境可搭配设置persistent_workers=True、pin_memory=True,避免每个batch重复初始化worker导致的阻塞。 - 数据集读取逻辑隐式卡住
先脱离Lightning单独测试数据集:实例化RCDataset后手动调用dataset[0],再循环取10个样本验证读取是否正常,重点检查是否存在路径拼接错误导致的IO等待、异常捕获逻辑吞掉了文件不存在/解码失败的报错、__getitem__里写了带全局锁的逻辑导致跨进程传输阻塞。 - 模型计算逻辑问题
注意PyTorch的nn.CrossEntropyLoss内置了LogSoftmax计算,如果在模型最后一层手动加了Softmax输出,虽然不会立刻抛错,但数值溢出产生NaN/inf时会导致迭代静默卡住,建议删掉模型末尾的Softmax层,直接输出原始logits给交叉熵损失计算。
可以单独生成3个和输入尺寸一致的随机Tensor,手动走一遍前向传播、损失计算、反向传播全流程,确认计算链路没有阻塞。 - 进度条组件假死
部分低版本PyTorch Lightning和tqdm适配有bug,会出现实际训练在跑但进度条不刷新的情况,可以在初始化Trainer时加参数enable_progress_bar=False关闭进度条,观察控制台是否正常输出step、epoch相关日志,确认是不是显示问题。
内容的提问来源于stack exchange,提问作者LuckyPotatoe
相关产品推荐
相关产品推荐

