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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.01 20:18:29