PyTorch Lightning升级后运行大幅变慢的原因排查与提速诉求
问题背景
升级前使用PyTorch Lightning 0.7.6 + Python 3.7,升级后为PyTorch Lightning 2.2.1 + Python 3.8,基于CUDA的PyTorch版本,但单轮epoch耗时从30分钟飙升至2.5小时,需在保留新框架特性的前提下恢复训练速度。
排查与优化方案
1. 修正设备管理逻辑
Lightning 2.x会自动处理模型的设备迁移,手动执行model.to(device)会触发不必要的设备同步开销,删除以下冗余代码:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model.to(device)
2. 优化训练器核心配置
- 启用混合精度训练:将
precision=32改为precision="16-mixed"(需CUDA支持自动混合精度AMP),大幅降低显存占用并提升计算效率:trainer = pl.Trainer( # ...其他参数 precision="16-mixed", ) - 明确训练策略:将
strategy='auto'改为strategy='ddp'(单GPU场景也可使用,比默认策略更高效),避免Lightning自动选择低效执行策略:strategy='ddp', - 缩减checkpoint保存量:
save_top_k=-1会保存所有epoch的checkpoint,带来巨大IO开销,改为保存Top N最优模型(比如save_top_k=3):checkpoint_callback = ModelCheckpoint( # ...其他参数 save_top_k=3, )
3. 数据加载性能优化
检查make_dataloaders的参数配置,补充以下优化项:
num_workers:设置为CPU核心数(如8或16),并行加载数据pin_memory=True:针对CUDA设备,减少数据从CPU到GPU的拷贝耗时
修改调用代码:
train_loader, val_loader = make_dataloaders( **conf['data'], **conf['training'], channels=slice(0, 4), num_workers=8, pin_memory=True )
4. 精准定位性能瓶颈
启用Lightning内置性能分析器,定位耗时模块:
from pytorch_lightning.profilers import PyTorchProfiler profiler = PyTorchProfiler( on_trace_ready=torch.profiler.tensorboard_trace_handler(os.path.join(exp_dir, "profiler")), schedule=torch.profiler.schedule(wait=1, warmup=1, active=3, repeat=2) ) trainer = pl.Trainer( # ...其他参数 profiler=profiler, )
运行后查看TensorBoard的Profiler标签,可直接找出数据加载、模型前向/反向传播中的耗时热点。
5. 细节优化
- 确认
torch.set_float32_matmul_precision('high')生效:该设置允许PyTorch使用TensorFloat32(TF32)加速矩阵乘法,需CUDA架构>=Ampere(如RTX 30/40系列)。 - 检查
BeamTasNetSystem的training_step/validation_step:避免不必要的.cpu()调用,尽量将计算逻辑留在GPU上,减少跨设备数据传输开销。
内容的提问来源于stack exchange,提问作者xuxuxuxu223
相关产品推荐
相关产品推荐

