基于PyTorch Lightning的多标签情感分析LSTM推理报错如何解决?
问题根因
PyTorch Lightning的trainer.predict()接口不支持直接传入原始Tensor作为输入,它要求接收的是PyTorch DataLoader实例、LightningDataModule实例,或者是返回批次数据的可迭代对象。你直接传预处理后的Tensor时,PL内部会尝试读取输入对象的batch_sampler属性来生成批次,Tensor本身没有这个属性,就会触发你遇到的报错。
修复步骤
1. 封装推理用数据集
先把预处理好的Tensor包装成PyTorch内置的TensorDataset,不需要自定义数据集类:
from torch.utils.data import TensorDataset, DataLoader # 根据你自己的编码输出调整入参,示例为输入是input_ids、attention_mask两个张量的场景 infer_dataset = TensorDataset(input_ids, attention_mask) # 仅单个输入张量的场景写法:infer_dataset = TensorDataset(processed_tensor)
2. 构造推理DataLoader
给数据集套上DataLoader,按需设置批次大小,推理不需要打乱数据顺序:
infer_dataloader = DataLoader(infer_dataset, batch_size=32, shuffle=False)
3. 调用predict接口
把构造好的DataLoader传给trainer.predict()即可:
predictions = trainer.predict(model, dataloaders=infer_dataloader)
输出的predictions是每个批次的预测结果组成的列表,你可以自行拼接成完整的结果张量。
可选单样本推理方案
如果是小批量/单样本快速测试,不想构造DataLoader,也可以直接调用模型的forward方法,注意开启eval模式、关闭梯度计算避免显存浪费:
import torch model.eval() with torch.no_grad(): # 给输入张量增加batch维度,模型默认接收批次输入 input_tensor = processed_tensor.unsqueeze(0) prediction = model(input_tensor)
内容的提问来源于stack exchange,提问作者Trip Trap
相关产品推荐
相关产品推荐

