LightningModule仅保留与训练/验证/测试强耦合的逻辑(前向传播、损失计算、优化器配置、指标更新),数据加载、日志路径、模型保存策略、超参解析、部署封装等必须移出;self.save_hyperparameters()须置于__init__开头且参数可序列化;数据增强不可塞入forward;自定义loss含非tensor操作需移至validation_step外部;enable_checkpointing和precision="16-mixed"在生产环境需谨慎启用;分布式训练须显式设strategy="ddp"和accelerator="gpu";max_epochs与max_steps二选一;logger=False后需手动添加Logger;部署时应导出剥离trainer属性的子模块self.net而非整个LightningModule;输入预处理需独立成函数;DataModule中setup()须显式初始化各stage数据集,耗时操作移入setup(),CombinedLoader需设persistent_workers=True;动态sampler状态需自行实现state_dict/load_state_dict。

LightningModule 里哪些逻辑必须放进去?哪些必须移出去?
Lightning 的核心约束不是“写得更少”,而是“职责分离更硬”。LightningModule 只能包含与模型训练/验证/测试强耦合的逻辑:前向传播、损失计算、优化器配置、指标更新。所有数据加载、日志路径、模型保存策略、超参解析、部署封装,都必须移出 LightningModule——否则一上线就会卡在 checkpoint 加载失败或 wandb 初始化冲突上。
-
self.save_hyperparameters()只能放在__init__开头,且传入的参数必须是可序列化的(不能是Dataset实例或 lambda 函数) - 数据增强逻辑别塞进
forward,它会被 JIT 编译或导出为 TorchScript,而torchvision.transforms中部分操作(如RandomResizedCrop)不支持 trace - 自定义 loss 如果含非 tensor 操作(比如调用
sklearn.metrics),必须移到validation_step外部计算,否则 DDP 下会因 rank 0 独占计算导致梯度同步异常
Trainer 配置项哪些能开,哪些开了反而坏事?
生产环境里最常误开的是 enable_checkpointing=True(默认开启)和 precision="16-mixed"。前者在无共享存储的多节点训练中,若没配 checkpoint_dir 为 NFS 路径,各 rank 会各自写 checkpoint 导致覆盖;后者在某些旧 GPU(如 P100)或混合 batch size 场景下,grad scaler 会 silently 失效,loss 突然 nan。
- 分布式训练必须显式设
strategy="ddp"(不要依赖自动推断),并确认accelerator="gpu"—— Lightning 2.0+ 在 CPU 环境下默认用cpustrategy,但若代码里写了.cuda()就会报Expected all tensors to be on the same device -
max_epochs和max_steps二选一即可,同时设会导致后者优先级更高,容易让早停(EarlyStopping)失效 - 用
logger=False关闭默认 logger 后,别忘了手动加TensorBoardLogger或WandbLogger,否则trainer.test()不输出指标
如何让 Lightning 模型真正可部署?
LightningModule 本身不是部署单元。直接 torch.jit.script(model) 会失败,因为 LightningModule 含大量 trainer runtime 属性(如 self.global_step)。真正要导出的是剥离后的 model 子模块。
- 在
configure_model或__init__里把网络主干单独实例化成self.net = MyBackbone(),导出时只对self.net调用torch.jit.trace() - 输入预处理必须从
DataModule拆出来,写成独立函数(如def preprocess(x: np.ndarray) -> torch.Tensor:),否则 TorchScript 无法 infer numpy → tensor 转换逻辑 - 用
torch.jit.save()保存后,部署端加载时别调LightningModule.load_from_checkpoint()—— 它依赖完整 trainer 环境,应改用torch.jit.load()+ 手动调用preprocess+net()
常见报错:AttributeError: 'LightningDataModule' object has no attribute 'train_dataloader'
这不是 DataModule 写错了,而是你调用了 trainer.fit(model, datamodule) 之后,又手动执行了 datamodule.train_dataloader()。Lightning 会在 fit 前自动调用该方法并缓存结果;若你在外部再调,且没重写 setup() 或漏了 self.train_ds 初始化,就会触发此错。
立即学习“Python免费学习笔记(深入)”;
- 确保
DataModule.setup(stage)中对每个 stage 显式初始化对应 dataset:if stage == "fit": self.train_ds = MyDataset(...) - 不要在
train_dataloader()里做耗时操作(如解压 zip、生成 cache 文件),它可能被反复调用;这类逻辑全移到setup() - 如果用
CombinedLoader或自定义 sampler,务必在train_dataloader()返回前加return DataLoader(..., persistent_workers=True),否则多 worker 下 epoch 间 worker 重启导致 seed 错乱
LightningDataModule 的 state_dict 和 load_state_dict 方法——它们默认为空,但如果你在训练中动态修改了 sampler 的 shuffle 状态或 batch sampler 的 epoch 计数,这些状态不会随 checkpoint 保存,恢复训练时数据流就可能错位。


















