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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 06:48:04