
在 PyTorch Lightning 中,不能直接在 on_train_epoch_end 回调中调用 trainer.predict(),因其会干扰训练状态;推荐改用手动批处理迁移 + predict_step 调用,并通过 transfer_batch_to_device 确保设备一致性。
在 pytorch lightning 中,不能直接在 `on_train_epoch_end` 回调中调用 `trainer.predict()`,因其会干扰训练状态;推荐改用手动批处理迁移 + `predict_step` 调用,并通过 `transfer_batch_to_device` 确保设备一致性。
在 Lightning 训练流程中,trainer.predict() 是一个完整预测循环入口,它会重置内部状态(如 _results)、切换模型模式(eval()/train())、管理设备分发与日志器连接等。当它被嵌套在 on_train_epoch_end 这类训练钩子中调用时,会与当前训练上下文(尤其是 LoggerConnector 所依赖的 _results 状态)发生冲突,从而触发 AssertionError: self.trainer._results is not None 错误——这正是你遇到的根本原因。
因此,最自然、轻量且符合 Lightning 设计哲学的方式,是绕过 trainer.predict(),直接在回调中复用模型自身的 predict_step 逻辑,但必须确保输入数据已正确迁移到目标设备(如 GPU、MPS 或 CPU)。Lightning 提供了标准方法 model.transfer_batch_to_device(batch, device, dataloader_idx) 来完成这一关键步骤,它能自动处理张量、嵌套结构(如字典、列表)及多设备兼容性(包括 MPS)。
以下是推荐的实现方式:
from lightning.pytorch.callbacks import Callback
import torch
class PlotCallback(Callback):
def on_train_epoch_end(self, trainer: L.Trainer, model: L.LightningModule) -> None:
# 获取预测数据加载器(可来自模型或自定义)
dataloader = model.predict_dataloader()
# 切换为评估模式(可选,但推荐显式声明)
model.eval()
with torch.no_grad(): # 禁用梯度,节省内存与计算
for batch in dataloader:
# ✅ 关键:将 batch 正确迁移到当前模型设备
batch = model.transfer_batch_to_device(batch, model.device, 0)
# ✅ 直接调用 predict_step,复用模型定义的推理逻辑
predictions = model.predict_step(batch)
# ? 此处可进行后处理、可视化、记录到 W&B 等
# 例如:plot_and_log_to_wandb(predictions, trainer.current_epoch)
# 可选:恢复训练模式(Lightning 通常自动处理,但显式更健壮)
model.train()
⚠️ 注意事项:
- 不要省略 model.eval() 和 torch.no_grad() —— 否则 BatchNorm/Dropout 行为异常,且可能意外积累梯度;
- transfer_batch_to_device 必须显式调用:即使 dataloader 已返回 GPU 张量,在多设备(如 MPS)或分布式场景下,模型 device 与 batch.device 仍可能不一致;
- 若 predict_dataloader 返回多个 dataloader(如验证集+测试集),需传入对应 dataloader_idx;
- 避免在回调中修改模型参数或优化器状态,保持回调的“只读”与“副作用可控”原则。
该方案完全规避了 trainer.predict() 的生命周期冲突,代码简洁、性能高效、设备鲁棒性强,是 Lightning 社区广泛采纳的惯用模式,也与官方文档中关于自定义推理逻辑的最佳实践高度一致。











