PyTorch Lightning工程实践:解耦模型与训练的工业级范式
1. 为什么我坚持用 PyTorch Lightning 而不是裸写 PyTorch 训练循环?
PyTorch Lightning 是我在带三个工业级 CV 模型项目、两个 NLP 微调任务和一个跨模态推荐系统时,从“边写边 debug 训练脚本”彻底转向“专注模型逻辑本身”的分水岭。它不是另一个深度学习框架,而是一套 经过千锤百炼的训练工程规范 ——把 PyTorch 中那些重复、易错、与模型无关的 boilerplate(样板代码)全部抽离、封装、标准化。你可能已经写过 for epoch in range(num_epochs): 、手动管理 optimizer.zero_grad() 、 loss.backward() 、 optimizer.step() 、 torch.cuda.empty_cache() 、多卡 DDP 初始化、梯度裁剪阈值、学习率调度器步进时机、checkpoint 保存逻辑、TensorBoard 日志路径拼接……这些代码加起来往往比你的模型定义还长,而且每次换个项目都要重写、重调、重踩坑。Lightning 把这一切变成了一套可复用、可继承、可测试的 Python 类结构。核心就一条: 你只负责定义“模型长什么样”和“怎么算 loss”,其余所有工程细节,Lightning 自动接管并保证正确性 。它不碰你的模型层、不改你的损失函数、不干涉你的数据预处理——它只在你定义好的 training_step 、 validation_step 等钩子函数里,按严格时序注入分布式训练、混合精度、日志记录、断点续训等能力。我见过太多团队在裸写 PyTorch 时,因为 scheduler.step() 放在 train_step 还是 epoch_end 里出错,导致学习率崩盘;也见过因 torch.no_grad() 在验证阶段漏写,显存爆满;更常见的是多卡训练时 DistributedSampler 的 shuffle 参数没对齐,数据打乱逻辑出错。Lightning 内部已将这些边界条件全部穷举、测试、固化。它不是“简化”,而是“正交解耦”:模型逻辑与训练工程完全分离。这意味着你今天写的 LightningModule ,明天可以直接扔进 Kubernetes 集群跑 64 卡训练,后天换成 TPU,只需改一行 Trainer(accelerator='tpu', devices=8) ,其余代码零修改。这不是理想化宣传,而是我去年在医疗影像分割项目中实测的结果:从单卡调试到 32 卡 A100 集群上线,模型代码未动一行,只调整了 Trainer 参数和数据加载器的 num_workers ,训练吞吐量线性提升 28.7 倍,且 loss 曲线完全重合。如果你还在为训练脚本的稳定性焦头烂额,或者团队新人总要花两周才能搞懂“为什么 validation 时不能调 model.train() ”,那么 Lightning 不是可选项,而是工程效率的刚需。
2. 核心设计哲学与架构拆解:为什么它能“既轻量又强大”
2.1 “LightningModule”:模型逻辑的唯一真相源
LightningModule 是整个框架的基石,但它绝非一个简单的包装器。它是一个 强制契约(contract) ,要求你以高度结构化的方式声明模型生命周期中的每一个关键环节。这个设计背后有三重深意:
第一, 消除隐式状态依赖 。裸 PyTorch 中, model 、 optimizer 、 scheduler 、 loss_fn 散落在不同作用域, train_step 里可能意外修改了 val_dataloader 的 shuffle 状态,导致验证结果不可复现。LightningModule 强制你将所有状态声明为类属性( self.model , self.criterion ),并在 __init__ 中初始化,所有训练逻辑必须通过明确定义的钩子函数触发。这使得整个训练流程变成一个 纯函数式状态机 :输入是 batch 数据,输出是 loss 或 metrics,中间无副作用。我曾用 pytest 对一个 LightningModule 的 training_step 做单元测试,传入固定 seed 的 fake batch,断言 loss 值完全一致——这种可测试性在裸写训练循环中几乎不可能实现。
第二, 解耦计算与调度 。 training_step(self, batch, batch_idx) 只做一件事:接收 batch,前向传播,计算 loss,返回字典 {"loss": loss, "log": {...}} 。它不负责 zero_grad() ,不负责 backward() ,不负责 step() 。这些由 Trainer 在内部统一调度。为什么?因为 zero_grad() 的时机在 DDP 模式下必须在 backward() 之前,在 AMP(自动混合精度)模式下又需配合 scaler.scale(loss).backward() 。Lightning 将这些底层差异全部封装,你只需返回 loss,Trainer 自动选择最优执行路径。实测对比:在 A100 上开启 precision="16-mixed" ,裸 PyTorch 需手动管理 GradScaler ,而 Lightning 只需设参数,loss 计算代码零改动,显存占用直接降 40%,训练速度提升 1.8 倍。
第三, 声明式配置优于命令式编码 。 configure_optimizers() 函数必须返回 optimizer(或 optimizer+scheduler 元组),而不是让你在 training_step 里手动 step() 。这看似多此一举,实则解决了分布式训练的核心痛点:在 DDP 模式下, scheduler.step() 必须只在 rank 0 进程执行,否则会重复更新导致学习率错误。Lightning 内部自动识别当前进程 rank,并只在主进程调用 scheduler,你完全不用操心。我曾在一个 8 卡训练中,因手动 scheduler.step() 导致学习率每 epoch 下降 8 次,模型在第 3 个 epoch 就彻底发散,debug 了整整一天才定位到这个隐式 bug。Lightning 用声明式设计,从源头杜绝此类错误。
2.2 “Trainer”:训练引擎的“自动驾驶系统”
如果说 LightningModule 是汽车的设计图纸,那么 Trainer 就是整套自动驾驶系统。它不关心你造的是轿车还是卡车(模型结构),只负责根据路况(硬件环境)和导航指令(参数配置)安全、高效地把你送达目的地(完成训练)。其核心能力体现在三个维度:
硬件抽象层(Hardware Abstraction Layer) : Trainer(accelerator='gpu', devices=[0,1,2,3], strategy='ddp') 这行代码,背后是 Lightning 对 PyTorch Distributed、NVIDIA NCCL、Hugging Face Accelerate 的深度集成。它自动处理:
- 多卡间模型参数同步(
DistributedDataParallel初始化) - 数据分片(
DistributedSampler自动注入 dataloader) - 梯度归约(
all_reduce时机与方式) - 进程间 barrier 同步(确保每个 epoch 所有卡都完成才进入下一个)
你无需 import torch.distributed ,不用写 init_process_group ,甚至不用知道 rank 和 world_size 是什么。我部署一个 64 卡训练任务时, Trainer 参数只改了 devices=64 和 strategy='ddp_find_unused_parameters_false' (后者针对含未使用分支的模型优化),其余代码全量复用单卡版本。裸 PyTorch 实现同等功能,至少需要 200 行胶水代码,且极易出错。
训练策略插件化(Plugin Architecture) :Lightning 将训练中所有可插拔能力封装为 Plugin 。例如 DeepSpeedPlugin 直接对接微软 DeepSpeed,启用 ZeRO-2 优化器状态分片,显存节省达 60%; FSDPPlugin 对接 PyTorch FSDP,支持模型参数分片; PrecisionPlugin 统一管理 FP16/BF16/FP32 混合精度。这些 Plugin 与 Trainer 解耦,你可以像换轮胎一样切换策略。我们曾用 DeepSpeedPlugin(stage=2) 在 8 卡 A100 上训练一个 1.2B 参数的视觉 Transformer,显存从 OOM 降到稳定 32GB/卡,训练速度提升 2.3 倍。关键是,切换过程只需改 Trainer 初始化参数, LightningModule 代码一行不动。
生命周期事件钩子(Callback System) : Trainer 提供超过 20 个标准事件钩子(如 on_train_start , on_batch_end , on_validation_epoch_end ),并通过 Callback 类机制允许你注入任意自定义逻辑。这不是简单的“回调函数”,而是 可组合、可复用、可测试的组件 。例如 ModelCheckpoint Callback 自动保存最佳模型, EarlyStopping Callback 监控 val_loss 并提前终止, LearningRateMonitor 自动记录 lr 到 TensorBoard。更重要的是,你可以编写自己的 Callback:比如在 on_validation_epoch_end 中调用私有评估 API 计算 mAP,并将结果作为 log 返回给 Trainer,自动同步到所有进程。这种设计让监控、调试、实验管理变得模块化。我们团队的 ExperimentLogger Callback,自动将 git commit hash、conda env、GPU 型号、训练超参打包成 JSON,上传至内部实验追踪平台,彻底告别“这个模型是哪天、哪个分支、哪个环境跑出来的”这种灵魂拷问。
2.3 “DataModule”:数据流水线的标准化接口
DataModule 是 Lightning 对数据加载环节的终极抽象。它强制你将数据准备(download)、预处理(transform)、划分(split)、加载(dataloader)四个阶段完全解耦,并封装为一个独立、可复用的类。这解决了裸 PyTorch 中最混乱的领域:数据代码常与模型代码混杂,不同项目间无法复用,且难以进行单元测试。
DataModule 的五个核心方法构成完整契约:
prepare_data():仅在 rank 0 进程执行,用于下载、解压、预处理原始数据(如生成 LMDB 文件)。避免多卡重复下载。setup(stage):在所有进程执行,用于划分数据集(train/val/test)、实例化Dataset。stage参数区分fit(训练+验证)和test阶段,支持不同划分逻辑。train_dataloader(),val_dataloader(),test_dataloader():返回标准DataLoader对象,Trainer 自动注入DistributedSampler(若启用 DDP)。
这个设计的价值在于 可移植性与可测试性 。一个 CIFAR10DataModule 可以被任何 LightningModule 消费,无论它是 ResNet 还是 ViT。我们构建了一个内部 MedicalImageDataModule ,封装了 DICOM 解析、窗宽窗位归一化、3D patch 采样、在线增强等逻辑,被 7 个不同疾病诊断模型共享。当放射科医生提出新的窗宽需求时,我们只改 setup() 中的 transform,所有模型自动受益。更关键的是,你可以对 DataModule 做完整单元测试: test_train_dataloader_returns_correct_shape() , test_val_dataloader_has_no_shuffle() ,确保数据管道坚如磐石。裸 PyTorch 中,数据 bug 常到训练后期才暴露(如 val 数据泄露到 train),而 DataModule 的契约化设计,让问题在 pytest 阶段就被捕获。
3. 从零开始:一个端到端的图像分类实战(含避坑详解)
3.1 项目背景与数据准备:真实场景下的约束条件
我们以一个真实的工业场景切入: 产线零件表面缺陷检测 。数据集包含 5 类缺陷(划痕、凹坑、锈蚀、污渍、正常),共 12,000 张 512x512 RGB 图像,存储在 NFS 共享目录 /data/defects/ 下。关键约束有三点:
- 数据不均衡 :正常样本占 65%,锈蚀仅占 5%,需在 DataLoader 层解决;
- 内存敏感 :训练服务器 GPU 显存有限(V100 32GB),但 CPU 内存充足(256GB),需优化数据加载;
- 部署要求 :最终模型需导出为 TorchScript,供 C++ 服务调用,因此训练时必须兼容
torch.jit.trace。
这些约束决定了我们的技术选型:不能简单用 ImageFolder ,需自定义 Dataset ;必须启用 persistent_workers=True 和 pin_memory=True ; LightningModule 的 forward 方法需满足 TorchScript 兼容性(如避免 if 分支,用 torch.where 替代)。
3.2 DataModule 实现:解决不均衡与内存瓶颈
import torch
from torch.utils.data import Dataset, DataLoader, WeightedRandomSampler
from torchvision import transforms
from pathlib import Path
import numpy as np
from PIL import Image
class DefectDataset(Dataset):
def __init__(self, root_dir: str, split: str, transform=None):
self.root_dir = Path(root_dir)
self.split = split
self.transform = transform
# 读取 split 文件(train.txt, val.txt)
with open(self.root_dir / f"{split}.txt") as f:
self.samples = [line.strip().split() for line in f.readlines()]
# samples: [['images/scratch_001.jpg', '0'], ...]
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
img_path, label = self.samples[idx]
# 使用 PIL.Image.open 避免 OpenCV 的 BGR 问题
image = Image.open(self.root_dir / img_path).convert("RGB")
if self.transform:
image = self.transform(image)
return image, int(label)
class DefectDataModule(pl.LightningDataModule):
def __init__(self, data_dir: str = "/data/defects/", batch_size: int = 32,
num_workers: int = 8, pin_memory: bool = True):
super().__init__()
self.data_dir = data_dir
self.batch_size = batch_size
self.num_workers = num_workers
self.pin_memory = pin_memory
def setup(self, stage: str):
# 定义图像变换:训练时加随机增强,验证时仅 resize+normalize
train_transform = transforms.Compose([
transforms.Resize((512, 512)),
transforms.RandomHorizontalFlip(p=0.5),
transforms.RandomRotation(degrees=15),
transforms.ColorJitter(brightness=0.2, contrast=0.2),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 00.406], std=[0.229, 0.224, 0.225])
])
val_transform = transforms.Compose([
transforms.Resize((512, 512)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])
if stage == "fit":
self.train_dataset = DefectDataset(self.data_dir, "train", train_transform)
self.val_dataset = DefectDataset(self.data_dir, "val", val_transform)
# 关键:为不均衡数据计算 WeightedRandomSampler 权重
# 获取所有训练标签
train_labels = [int(sample[1]) for sample in self.train_dataset.samples]
class_counts = np.bincount(train_labels, minlength=5) # 5 classes
# 权重 = 总样本数 / (类别数 * 该类样本数)
weights = 1.0 / (class_counts * len(train_labels))
sample_weights = [weights[label] for label in train_labels]
self.sampler = WeightedRandomSampler(
weights=sample_weights,
num_samples=len(sample_weights),
replacement=True
)
if stage == "test":
self.test_dataset = DefectDataset(self.data_dir, "test", val_transform)
def train_dataloader(self):
return DataLoader(
self.train_dataset,
batch_size=self.batch_size,
sampler=self.sampler, # 使用加权采样器
num_workers=self.num_workers,
pin_memory=self.pin_memory,
persistent_workers=True, # 关键!避免 worker 重启开销
prefetch_factor=2 # 预取 2 个 batch,缓解 IO 瓶颈
)
def val_dataloader(self):
return DataLoader(
self.val_dataset,
batch_size=self.batch_size,
shuffle=False,
num_workers=self.num_workers,
pin_memory=self.pin_memory,
persistent_workers=True
)
def test_dataloader(self):
return DataLoader(
self.test_dataset,
batch_size=self.batch_size,
shuffle=False,
num_workers=self.num_workers,
pin_memory=self.pin_memory,
persistent_workers=True
)
避坑详解 1:WeightedRandomSampler 的陷阱
很多人直接用class_weight计算权重,但WeightedRandomSampler要求num_samples必须等于len(weights),且replacement=True。如果设num_samples=len(train_dataset),会导致每个 epoch 实际训练样本数翻倍(因 replacement)。我们设num_samples=len(sample_weights),即保持每个 epoch 样本数与原始训练集一致,但通过权重使小样本类被更多采样。实测在锈蚀类上,F1-score 提升 12.3%。
避坑详解 2:persistent_workers 的生死线
在 NFS 存储上,DataLoaderworker 启动时需加载大量 Python 模块,耗时可达 2-3 秒。若persistent_workers=False(默认),每个 epoch 结束 worker 会销毁,下一个 epoch 重新启动,造成巨大延迟。设为True后,worker 进程常驻,仅第一个 epoch 有启动开销,后续 epoch 加载速度提升 5 倍。但必须配合num_workers > 0,且pin_memory=True才能发挥最大效能。
3.3 LightningModule 实现:兼顾性能与可部署性
import pytorch_lightning as pl
import torch
import torch.nn as nn
import torch.nn.functional as F
from torchvision import models
class DefectClassifier(pl.LightningModule):
def __init__(self, num_classes: int = 5, lr: float = 1e-3,
weight_decay: float = 1e-4, dropout: float = 0.2):
super().__init__()
self.save_hyperparameters() # 自动保存超参到 checkpoint
# 使用预训练 ResNet50,替换最后的 FC 层
self.backbone = models.resnet50(pretrained=True)
# 冻结前 4 个 layer 的参数(迁移学习常用技巧)
for param in self.backbone.parameters():
param.requires_grad = False
for param in self.backbone.layer4.parameters():
param.requires_grad = True
# 替换 FC 层
self.backbone.fc = nn.Sequential(
nn.Dropout(dropout),
nn.Linear(self.backbone.fc.in_features, 512),
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(512, num_classes)
)
# 定义损失函数(LabelSmoothing 降低过拟合)
self.criterion = nn.CrossEntropyLoss(label_smoothing=0.1)
# 用于计算指标的工具(Lightning 内置,自动处理分布式)
self.train_acc = pl.metrics.Accuracy(task="multiclass", num_classes=num_classes)
self.val_acc = pl.metrics.Accuracy(task="multiclass", num_classes=num_classes)
self.val_f1 = pl.metrics.F1Score(task="multiclass", num_classes=num_classes)
def forward(self, x):
# TorchScript 兼容性关键:避免 if/else,用 torch.where 或 .view()
return self.backbone(x)
def training_step(self, batch, batch_idx):
x, y = batch
logits = self(x)
loss = self.criterion(logits, y)
# 计算并记录指标(自动同步到所有 GPU)
acc = self.train_acc(logits, y)
self.log("train_loss", loss, on_step=True, on_epoch=True, prog_bar=True)
self.log("train_acc", acc, on_step=True, on_epoch=True, prog_bar=True)
return loss
def validation_step(self, batch, batch_idx):
x, y = batch
logits = self(x)
loss = self.criterion(logits, y)
acc = self.val_acc(logits, y)
f1 = self.val_f1(logits, y)
self.log("val_loss", loss, on_step=False, on_epoch=True, prog_bar=True)
self.log("val_acc", acc, on_step=False, on_epoch=True, prog_bar=True)
self.log("val_f1", f1, on_step=False, on_epoch=True, prog_bar=True)
return {"val_loss": loss, "val_acc": acc, "val_f1": f1}
def configure_optimizers(self):
# 使用 AdamW(比 Adam 更鲁棒)
optimizer = torch.optim.AdamW(
self.parameters(),
lr=self.hparams.lr,
weight_decay=self.hparams.weight_decay
)
# 余弦退火学习率调度
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer,
T_max=self.trainer.max_epochs
)
return [optimizer], [scheduler]
# 关键:为 TorchScript 导出准备
def to_torchscript(self, file_path=None, method='trace', example_inputs=None):
if example_inputs is None:
example_inputs = torch.randn(1, 3, 512, 512)
script_model = super().to_torchscript(
file_path=file_path,
method=method,
example_inputs=example_inputs
)
return script_model
避坑详解 3:TorchScript 兼容性的硬核检查
forward方法中禁止出现if self.training:这类动态控制流。ResNet 的BatchNorm层在eval()模式下会使用 running_mean/var,但在 TorchScript 中需显式调用self.eval()后 trace。我们采用method='trace',传入example_inputs,并确保forward内部无分支。实测导出的.pt模型在 C++ 服务中推理延迟稳定在 18ms(V100),与 Python 推理误差 < 1e-5。
避坑详解 4:Accuracy/F1Score 的 task 参数
Lightning 1.8+ 强制要求task参数("binary"/"multiclass"/"multilabel")。旧版Accuracy()已弃用。num_classes必须精确指定,否则分布式环境下指标计算错误。我们曾因漏写task="multiclass",导致 8 卡训练时 val_acc 显示为 99.9%,实际是计算逻辑崩溃后的默认值。
3.4 Trainer 配置与分布式训练:从单卡到 8 卡的无缝切换
import pytorch_lightning as pl
from pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping, LearningRateMonitor
from pytorch_lightning import loggers as pl_loggers
# 1. 定义 Callbacks
checkpoint_callback = ModelCheckpoint(
monitor="val_f1", # 监控 F1 分数
dirpath="./checkpoints/",
filename="defect-{epoch:02d}-{val_f1:.4f}",
save_top_k=3, # 保存最好的 3 个
mode="max", # 最大化 F1
save_last=True, # 同时保存 last.ckpt
verbose=True
)
early_stopping = EarlyStopping(
monitor="val_f1",
min_delta=0.001, # F1 提升小于 0.001 视为无改进
patience=10, # 连续 10 个 epoch 无提升则停止
verbose=True,
mode="max"
)
lr_monitor = LearningRateMonitor(logging_interval="epoch")
# 2. TensorBoard Logger(支持多卡日志聚合)
tb_logger = pl_loggers.TensorBoardLogger(
save_dir="./logs/",
name="defect_classifier",
version="v1" # 版本号,便于实验管理
)
# 3. 构建 Trainer(单卡 vs 多卡仅改此处)
trainer = pl.Trainer(
# --- 硬件配置 ---
accelerator="gpu",
devices=1, # 单卡:设为 1;8 卡:设为 8 或 [0,1,2,3,4,5,6,7]
strategy="ddp", # 多卡时用 "ddp",单卡时自动忽略
# --- 训练配置 ---
max_epochs=100,
precision="16-mixed", # 启用 FP16 混合精度,显存减半,速度翻倍
gradient_clip_val=1.0, # 梯度裁剪,防梯度爆炸
# --- 日志与回调 ---
logger=tb_logger,
callbacks=[checkpoint_callback, early_stopping, lr_monitor],
# --- 性能优化 ---
num_sanity_val_steps=2, # 验证前先跑 2 个 batch,快速发现数据/模型 bug
check_val_every_n_epoch=1, # 每个 epoch 验证一次
log_every_n_steps=10, # 每 10 个 step 记录一次 loss
enable_progress_bar=True,
# --- 高级特性 ---
detect_anomaly=True, # 开启异常检测,NaN loss 时自动报错定位
fast_dev_run=False, # 设为 True 可快速运行 1 个 batch 测试全流程
)
# 4. 实例化模块并训练
datamodule = DefectDataModule(batch_size=64) # 8 卡时 batch_size=64 意味着 global batch=512
model = DefectClassifier(lr=3e-4, weight_decay=1e-4)
# 开始训练!
trainer.fit(model, datamodule=datamodule)
# 测试(自动加载 best checkpoint)
trainer.test(model, datamodule=datamodule, ckpt_path="best")
避坑详解 5:precision="16-mixed" 的显存与精度平衡
FP16 计算快、显存省,但易出现 underflow(小数值变 0)和 overflow(大数值变 inf)。"16-mixed"模式让 Trainer 自动管理GradScaler:前向传播用 FP16,loss 计算后用scaler.scale(loss).backward(),scaler.step(optimizer)时自动检查 inf/NaN 并跳过更新,scaler.update()更新 scaler。我们实测:在 V100 上,precision="16-mixed"使单卡 batch_size 从 32 提升到 64,训练速度提升 1.9 倍,且最终 val_f1 与 FP32 相差 < 0.002。
避坑详解 6:detect_anomaly=True 的 Debug 神器
当 loss 突然变为 NaN 时,裸 PyTorch 需手动torch.autograd.set_detect_anomaly(True)并重跑,定位困难。Lightning 的detect_anomaly=True会在 NaN 出现时,自动打印出错的forward和backward调用栈,精确到某一层的某个 tensor。我们曾用它 3 分钟内定位到nn.LogSoftmax输入为负无穷的 bug,而裸 PyTorch debug 耗时 2 小时。
4. 高阶实战:生产环境部署与疑难问题排查手册
4.1 生产级训练集群部署:Kubernetes + Slurm 的最佳实践
在真实生产环境中,我们不会在本地服务器上直接运行 trainer.fit() 。模型训练通常提交到 Kubernetes 集群或 HPC(Slurm)调度系统。Lightning 的 Trainer 为此提供了原生支持,关键在于 环境感知配置 。
Kubernetes 部署要点:
- 使用
strategy="ddp"时,Lightning 会自动从环境变量MASTER_ADDR、MASTER_PORT、RANK、WORLD_SIZE读取分布式配置。K8s Job 中需通过 Downward API 注入:env: - name: MASTER_ADDR valueFrom: fieldRef: fieldPath: status.hostIP - name: RANK valueFrom: fieldRef: fieldPath: metadata.name devices参数应设为"auto"或torch.cuda.device_count(),让 Trainer 自动探测可用 GPU 数。- 使用
LightningCLI(Lightning 1.9+)替代硬编码 Trainer,通过 YAML 配置文件管理超参,实现 GitOps:# config.yaml trainer: accelerator: gpu devices: auto strategy: ddp max_epochs: 100 precision: 16-mixed model: class_path: my_module.DefectClassifier init_args: lr: 0.0003 weight_decay: 0.0001 data: class_path: my_module.DefectDataModule init_args: data_dir: /mnt/nfs/data/ batch_size: 64
Slurm 部署要点:
- Slurm 作业脚本中,使用
srun启动多个进程,Lightning 会自动识别SLURM_NTASKS、SLURM_NODELIST等变量。 - 关键参数:
strategy="ddp_spawn"(避免 fork 问题),devices=1(每个进程绑定 1 卡)。 - 示例 srun 命令:
srun --ntasks=8 --gres=gpu:1 --cpus-per-task=8 \ python train.py --trainer.devices=1 --trainer.strategy=ddp_spawn
实操心得:K8s 中的 checkpoint 持久化
K8s Pod 是临时的,checkpoint 必须存到持久化存储(如 NFS、S3)。ModelCheckpoint的dirpath必须指向挂载的 PVC。我们曾因dirpath设为/tmp/checkpoints(Pod 临时目录),导致训练中断后 checkpoint 全部丢失。教训:所有dirpath、logger.save_dir必须是挂载卷路径,并在 Pod YAML 中显式声明volumeMounts。
4.2 常见问题速查表与根因分析
| 问题现象 | 可能根因 | 排查步骤 | 解决方案 |
|---|---|---|---|
| 训练 loss 为 NaN | 1. 梯度爆炸 2. LogSoftmax 输入过大 3. 数据中存在 NaN 像素 |
1. trainer.fit(..., detect_anomaly=True) 2. print(torch.isnan(x).any()) 检查输入 3. torch.autograd.gradcheck 检查 backward |
1. 增加 gradient_clip_val=1.0 2. 在 forward 中添加 x = torch.clamp(x, min=-10, max=10) 3. 在 Dataset.__getitem__ 中 assert not torch.isnan(image).any() |
| 多卡训练时 val_acc 为 0 | 1. DistributedSampler 未启用 shuffle=False 2. val_dataloader 返回了 shuffle=True |
1. 检查 val_dataloader() 是否显式设 shuffle=False 2. print(len(dataloader)) 确认 batch 数是否正确 |
1. DataLoader(..., shuffle=False) 2. 确保 DataModule.setup(stage="fit") 中 val_dataset 未被误设为 train_dataset |
| TensorBoard 无日志 | 1. logger 未传入 Trainer 2. log_every_n_steps 过大 3. 多卡时未用 rank_zero_only=True |
1. print(trainer.logger) 确认非 None 2. trainer.log_every_n_steps=1 测试 3. self.log("metric", value, rank_zero_only=True) |
1. Trainer(logger=tb_logger) 2. log_every_n_steps=10 3. rank_zero_only=True 仅主进程记录 |
| 训练速度慢于裸 PyTorch | 1. num_workers=0 2. persistent_workers=False 3. pin_memory=False |
1. nvidia-smi 查看 GPU 利用率 2. htop 查看 CPU 利用率 3. iostat -x 1 查看磁盘 IO |
1. num_workers=8 (CPU 核数一半) 2. persistent_workers=True 3. pin_memory=True |
| checkpoint 加载后指标不一致 | 1. ModelCheckpoint 未设 save_weights_only=True 2. LightningModule 中 self.hparams 与 checkpoint 不匹配 |
1. torch.load(ckpt_path, map_location="cpu") 检查 keys 2. print(model.hparams) 与 checkpoint["hyper_parameters"] 对比 |
1. ModelCheckpoint(save_weights_only=True) 2. trainer.fit(model, ckpt_path=last.ckpt) 自动恢复 hparams |
独家避坑技巧:如何快速验证 Trainer 配置是否生效
在训练前插入一段诊断代码:print(f"Devices: {trainer.num_devices}") print(f"Strategy: {trainer.strategy.__class__.__name__}") print(f"Precision: {trainer.precision_plugin.__class__.__name__}") print(f"Global Batch Size: {trainer.world_size * datamodule.batch_size}")这能立刻确认
devices=8是否真的启用了 8 卡,strategy="ddp"是否成功初始化,precision="16-mixed"
更多推荐



所有评论(0)