Python编写for循环遍历列表实现ELECTRA模型批量多标签预测
ELECTRA多标签分类模型批量预测实现方案
核心逻辑
批量预测本质是把单条预测的流程复用到全量数据集上,千条级小数据量可以直接逐行推理,万条级以上大数据量建议用PyTorch自带的DataLoader做分块批量推理,提升显存利用率和推理速度。
注意:推理前必须将模型切换为评估模式,同时关闭梯度计算,避免Dropout等训练层影响预测结果,同时减少不必要的显存占用。
方案1:小数据集简易实现(千条级样本适用)
逻辑简单无需额外封装,直接遍历数据集逐行预测后拼接结果即可,代码如下:
import pandas as pd import torch # 加载待预测数据集,csv格式用read_csv,tsv格式添加参数sep='\t' df = pd.read_csv("your_dataset_filepath.csv") # 替换为你的数据集实际路径 trained_model.eval() all_predictions = [] # 关闭梯度计算 with torch.no_grad(): for text in df["sample_text"].values: # 文本编码 encoding = tokenizer.encode_plus( text, add_special_tokens=True, max_length=512, return_token_type_ids=False, padding="max_length", truncation=True, # 补充截断逻辑,避免超长文本报错 return_attention_mask=True, return_tensors='pt', ) # 模型推理 _, pred_res = trained_model(encoding["input_ids"], encoding["attention_mask"]) all_predictions.append(pred_res.flatten().numpy()) # 将预测结果转为DataFrame,列名与训练时的标签列保持一致 pred_df = pd.DataFrame(all_predictions, columns=LABEL_COLUMNS) # 预测结果与原数据集横向拼接 final_result = pd.concat([df, pred_df], axis=1) # 保存结果,utf-8-sig编码避免Excel打开乱码 final_result.to_csv("batch_prediction_result.csv", index=False, encoding="utf-8-sig")
方案2:大数据集批量加速实现(万条级以上样本适用)
通过DataLoader实现分批次推理,显存利用率更高,支持GPU加速,速度比逐行推理快5~20倍,代码如下:
import pandas as pd import torch from torch.utils.data import Dataset, DataLoader # 自定义文本数据集类 class MultiLabelTextDataset(Dataset): def __init__(self, text_list, tokenizer, max_seq_len=512): self.text_list = text_list self.tokenizer = tokenizer self.max_seq_len = max_seq_len def __len__(self): return len(self.text_list) def __getitem__(self, index): text = str(self.text_list[index]) encoding = self.tokenizer.encode_plus( text, add_special_tokens=True, max_length=self.max_seq_len, return_token_type_ids=False, padding="max_length", truncation=True, return_attention_mask=True, return_tensors='pt', ) return { "input_ids": encoding["input_ids"].flatten(), "attention_mask": encoding["attention_mask"].flatten() } # 1. 加载待预测数据集 df = pd.read_csv("your_dataset_filepath.csv") # 2. 初始化数据集与加载器,batch_size根据显存大小调整,可选值8/16/32 predict_dataset = MultiLabelTextDataset(df["sample_text"].values, tokenizer) predict_loader = DataLoader(predict_dataset, batch_size=16, shuffle=False) trained_model.eval() all_predictions = [] with torch.no_grad(): for batch in predict_loader: input_ids = batch["input_ids"] attention_mask = batch["attention_mask"] # 如果有GPU,取消下面两行注释将张量移到GPU,推理速度大幅提升 # input_ids = input_ids.cuda() # attention_mask = attention_mask.cuda() _, batch_pred = trained_model(input_ids, attention_mask) all_predictions.extend(batch_pred.numpy()) # 结果拼接与保存逻辑和简易方案一致 pred_df = pd.DataFrame(all_predictions, columns=LABEL_COLUMNS) final_result = pd.concat([df, pred_df], axis=1) final_result.to_csv("batch_prediction_result.csv", index=False, encoding="utf-8-sig")
注意事项
- 原单条预测代码缺少
truncation=True参数,遇到长度超过512的文本会触发报错,必须补上 - GPU推理时需要先把模型也加载到GPU上:
trained_model = trained_model.cuda() batch_size不要设置过大,否则会触发显存溢出,建议从16开始测试,根据显存余量上下调整- 输出结果会保留所有标签列的预测概率,和期望输出结构完全一致,如果需要0/1二值结果,加一步阈值判断即可
内容的提问来源于stack exchange,提问作者Quantizer
相关产品推荐
相关产品推荐

