如何使用Huggingface的DistilBERT模型完成新数据情感分析及标签预测
你调用model.predict()报错是因为原生PyTorch版DistilBERT序列分类模型没有内置predict方法,该方法是Huggingface Trainer类的专属方法,你可以根据自己的训练方式选择对应预测逻辑,具体实现如下:
前置步骤:处理新数据
首先对DataFrame里的待预测文本做编码,参数要和训练时保持一致:
import torch import pandas as pd from transformers import DistilBertTokenizerFast # 替换成你的DataFrame和对应的文本列名 df = pd.read_csv("你的待预测数据文件.csv") new_texts = df["review"].tolist() # 复用训练时用的分词器 tokenizer = DistilBertTokenizerFast.from_pretrained("distilbert-base-uncased") new_encodings = tokenizer(new_texts, truncation=True, padding=True, return_tensors="pt")
方案1:使用Trainer训练的模型预测
如果你之前是用Trainer完成的训练,直接调用Trainer的predict方法即可:
# 复用你之前定义的IMDbDataset类,新数据没有标签可以传占位值 class IMDbDataset(torch.utils.data.Dataset): def __init__(self, encodings, labels): self.encodings = encodings self.labels = labels def __getitem__(self, idx): item = {key: torch.tensor(val[idx]) for key, val in self.encodings.items()} item['labels'] = torch.tensor(self.labels[idx]) return item def __len__(self): return len(self.labels) # 构造预测数据集,标签随便填占位值即可,预测时不会用到 dummy_labels = [0] * len(new_texts) new_dataset = IMDbDataset(new_encodings, dummy_labels) # 调用训练好的trainer做预测 predict_output = trainer.predict(new_dataset) # 取出模型输出的logits,转成预测标签 pred_labels = predict_output.predictions.argmax(axis=1).tolist() # 把预测结果写入原DataFrame df["pred_label"] = pred_labels df["sentiment"] = df["pred_label"].map({1: "积极", 0: "消极"})
方案2:使用原生PyTorch训练的模型预测
如果你是用原生PyTorch手写的训练循环,按如下逻辑实现推理即可:
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu") # 把模型切到评估模式,关闭dropout、层归一化的训练态逻辑 model.to(device) model.eval() # 关闭梯度计算,加快推理速度、降低显存占用 with torch.no_grad(): input_ids = new_encodings["input_ids"].to(device) attention_mask = new_encodings["attention_mask"].to(device) outputs = model(input_ids, attention_mask=attention_mask) # 取概率最高的类别作为预测结果 pred_labels = outputs.logits.argmax(dim=1).cpu().tolist() # 写入预测结果 df["pred_label"] = pred_labels df["sentiment"] = df["pred_label"].map({1: "积极", 0: "消极"})
小提示
如果待预测数据量很大,可以拆分成小批量分批推理,避免显存溢出。
内容的提问来源于stack exchange,提问作者brownie_coder
相关产品推荐
相关产品推荐

