Transformers pipeline 怎么批量分类文本?保留行号并核对输出

用现成英文分类 pipeline 分组推理,保存原始行号并核对输入输出条数;说明 CPU 批大小、超长截断与得分的解释边界。

需要把一组英文文本交给现成分类模型时,可以用 Transformers pipeline 统一分词、推理和输出格式。批量处理的关键是保留原始行号、检查每条输入对应一条输出,并且把推理得分与业务准确率分开。下面演示英文情感二分类,不训练模型。

准备明确的模型和短样本

在独立 Python 环境安装 transformers==4.57.1 和适合设备的 PyTorch。CPU 用户可先执行 python -m pip install torch==2.6.0 --index-url https://download.pytorch.org/whl/cpu。首次运行会下载模型权重和分词文件,需要网络及足够的磁盘和内存。

Transformers pipeline 怎么批量分类文本?保留行号并核对输出

模型卡说明此检查点来自 SST-2 的英文分类任务,许可为 Apache-2.0,并提醒存在偏差。本题不把它推荐为中文通用分类器。以下句子是人为编写的演示输入;本次没有运行这个检查点的实际推理或速度评测。

按小组推理并写出行号

import json
from pathlib import Path
from transformers import pipeline

rows = [
    {"row_id": 101, "text": "The answer was helpful."},
    {"row_id": 102, "text": "The reply was confusing."},
    {"row_id": 103, "text": "The issue was solved quickly."},
]
checkpoint = "distilbert/distilbert-base-uncased-finetuned-sst-2-english"
classifier = pipeline(
    "text-classification", model=checkpoint, tokenizer=checkpoint,
    framework="pt", device=-1,
)
results = []
for start in range(0, len(rows), 2):
    group = rows[start:start + 2]
    texts = [row["text"] for row in group]
    if any(not text.strip() for text in texts):
        raise ValueError("本组含空文本,先按行号核对")
    predictions = classifier(
        texts, batch_size=1, truncation=True, max_length=128
    )
    if len(predictions) != len(group):
        raise RuntimeError("输入输出条数不一致,停止写结果")
    for row, prediction in zip(group, predictions):
        results.append({
            "row_id": row["row_id"], "text": row["text"],
            "label": prediction["label"], "score": float(prediction["score"]),
        })

assert [r["row_id"] for r in results] == [r["row_id"] for r in rows]
Path("sentiment-demo.json").write_text(
    json.dumps(results, ensure_ascii=False, indent=2), encoding="utf-8"
)
print("完成条数:", len(results))
print(json.dumps(results, ensure_ascii=False, indent=2))

Pipeline API支持给文本分类流水线传入字符串列表,并通过 batch_size 控制模型推理批大小。示例一次组织两条记录,但在 CPU 上让模型逐条前向:应用分组大小和模型 batch 大小是两个不同参数。

从哪些结果判断处理链跑通

验证方法是先看到完成条数为 3,结果文件含 101、102、103 且顺序不变,每条有 label 与 score;然后人工查看原文本是否与标签对应。该模型配置应使用 POSITIVE 和 NEGATIVE 标签,得分在 0 到 1 之间;具体标签和数值由实际运行产生,本文没有编造预测值。

有输出不等于分类可靠。先准备人工标注的代表性验证集,再统计误报与漏报;score 是模型输出得分,不能直接写成“这条结论有同等概率正确”。对否定、讽刺、领域词和中文文本尤其需要单独核验。

批大小、截断和大文件的边界

本例的 max_length=128 约束的是 token 长度而非汉字数,超长文本会被截断;若判断依据在末尾,标签可能受影响。先检查长度分布,决定分段还是换支持长度更合适的模型,不能把“没报错”当作全文已读完。

官方文档提醒 batching 不总能提高速度,CPU 可从 batch_size=1 开始。GPU 场景也要先用代表性长度的小样本测量,再逐渐增加;长度混杂会增加填充,可能使速度变慢或内存不足。本教程不承诺加速倍数。

换成 CSV 大文件时应分块读取和分块保存,逐块核对行号;本文的 results 列表只适合小演示。发生下载、权限、内存或结果结构错误时停止本组写入,保留失败行号,修好后只重做该组。完成本例后,读者下一步是替换为一小组获准使用且人工标注的英文文本,而不是直接全量上线。

Ai菜鸟网。发布者:AI小管家,转载请注明出处:https://www.alyyhw.com/30442.html

赞 (0)
AI小管家的头像AI小管家
Transformers 聊天模板怎么用?把角色消息转成模型输入
上一篇 7小时前
Transformers 批量文本长度不一致怎么办?核对填充、截断与 attention_mask
下一篇 7小时前

相关推荐

联系我们

联系我们

1

在线咨询: QQ交谈

邮件:admin@example.com

工作时间:周一至周五,9:30-18:30,节假日休息

关注微信
关注微信
分享本页
返回顶部