You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于PyTorch-Lightning与Swin Transformer的模型预测报错问题求助

解决PyTorch-Lightning模型预测的报错问题

让我们一步步拆解你遇到的问题,然后给出最贴合Lightning规范的解决方案:

错误原因分析

  1. 第一个错误ModuleAttributeError: 'Model' object has no attribute 'predict':你的自定义Model类继承自LightningModule,但并没有实现predict方法——LightningModule默认不带这个方法,要么自己实现,要么用官方推荐的Trainer.predict流程。
  2. 第二个错误AttributeError: 'PetfinderDataModule' object has no attribute 'shape':你直接把整个DataModule对象传给了model,但model.forward()接收的是张量输入(单批图像数据),不是数据集对象,自然会报错。

方法一:用PyTorch Lightning Trainer做预测(最规范推荐)

这是Lightning官方设计的预测流程,能自动处理设备转移、数据加载、梯度关闭等细节:

  1. 正确加载模型(替代手动load_state_dict)
# 用Lightning的load_from_checkpoint直接从ckpt恢复完整模型,自动处理超参数和权重
model = Model.load_from_checkpoint(
    checkpoint_path=f'{config.model.name}/default/version_0/checkpoints/best_loss.ckpt',
    cfg=config
)
model = model.cuda().eval()
  1. 实例化Trainer并执行预测
from pytorch_lightning import Trainer

# 根据你的硬件选择accelerator:'gpu'或'cpu',devices指定GPU数量
trainer = Trainer(accelerator='gpu', devices=1)
# 调用trainer.predict,传入模型和DataModule
test_predictions = trainer.predict(model, datamodule=test_dataset)

# 把各batch的预测结果拼接成完整张量
test_predictions = torch.cat(test_predictions)

方法二:手动遍历数据加载器预测(适合轻量场景)

如果你不想依赖Trainer,可以手动获取测试集的DataLoader,逐批处理:

  1. 获取测试数据加载器
# 从DataModule中取出测试集的DataLoader
test_dataloader = test_dataset.test_dataloader()
  1. 逐批执行推理
model.eval()
test_predictions = []

# 推理阶段关闭梯度计算,节省内存和计算资源
with torch.no_grad():
    for batch in test_dataloader:
        images, _ = batch  # 你的batch结构是(图像, 标签),推理时可以忽略标签
        images = images.cuda()
        
        # 注意:你当前在模型内部用的是训练时的transform,建议替换为测试专用变换(比如去掉随机增强)
        # 后续可以把变换逻辑移到DataModule中,分训练/测试分别处理
        images = model.transform(images)
        
        # 调用forward得到预测结果,和训练时的后处理保持一致
        logits = model(images).squeeze(1)
        preds = logits.sigmoid().detach().cpu() * 100.
        test_predictions.append(preds)

# 拼接所有batch的结果
test_predictions = torch.cat(test_predictions)

额外优化建议

  • 数据变换的规范处理:你现在在模型的__share_step里应用训练变换,这会导致推理时也用到随机裁剪、翻转等训练增强操作,影响预测结果。建议把变换逻辑移到PetfinderDataModule中,在setup方法里分别定义训练和测试的专属变换。
  • 模型加载的最佳实践:pl.load_from_checkpoint比手动load_state_dict更可靠,能自动恢复模型的超参数、权重等所有状态,避免手动处理state_dict的潜在问题。

内容的提问来源于stack exchange,提问作者Khadija

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.30 23:29:10