如何使用Python中的Transformers库进行文本分类?

雨丽同学_3801

雨丽同学_3801

2026-09-10

817人浏览

原创

加载预训练模型和分词器必须使用同一模型名,否则因词汇表、特殊token及输入结构不一致导致indexerror或logits维度错误;输入需经tokenizer转为pytorch张量并设truncation=true、padding=true;预测须先softmax再argmax,并用model.config.id2label映射标签;trainer需显式传compute_metrics函数才能输出准确率等指标。

如何使用python中的transformers库进行文本分类?

加载预训练模型和分词器时,AutoModelForSequenceClassificationAutoTokenizer 必须匹配同一模型名

直接用错模型名是新手最常踩的坑——比如用 "bert-base-uncased" 加载分词器,却用 "roberta-base" 加载模型。两者 tokenizer 的词汇表、特殊 token(如 [CLS]<s></s>)和模型输入结构不一致,会导致 forward()IndexError: index out of range in self 或 logits 维度错乱。

正确做法是共用一个模型标识符:

from transformers import AutoTokenizer, AutoModelForSequenceClassification

model_name = "distilbert-base-uncased-finetuned-sst-2-english"
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForSequenceClassification.from_pretrained(model_name)

注意:model_name 可以是 Hugging Face Hub 上的公开 checkpoint,也可以是你本地微调后保存的路径(如 "./my_finetuned_model"),但必须确保该路径下同时存在 config.jsonpytorch_model.bintokenizer_config.json 等文件。

输入文本需经 tokenizer(..., return_tensors="pt") 转为 PyTorch 张量,不能直接传字符串

模型的 forward() 接口只接受 input_idsattention_mask 这类张量,传入原始字符串会报 TypeError: expected Tensor as element 0 in argument 0, but got str

关键参数不能漏:

  • return_tensors="pt":强制返回 PyTorch tensor(不是 Python list 或 NumPy array)
  • truncation=True:超长文本必须截断,否则可能触发 CUDA OOM 或 shape mismatch
  • padding=True:批量推理时保证 batch 内序列等长;单条推理可省略,但加了更稳妥

示例:

inputs = tokenizer(
    "This movie is terrible.", 
    return_tensors="pt", 
    truncation=True, 
    padding=True
)
outputs = model(**inputs)
logits = outputs.logits

获取预测标签要先过 torch.nn.functional.softmax,再用 torch.argmax

模型输出的 logits 是未归一化的分数,直接取 argmax 容易误判(尤其多分类时类别置信度接近)。必须先 softmax 得到概率分布,再取最大概率索引。

python全能编程助手
python全能编程助手

SkillSub Pro - Python 题解与代码注释双功能技能功能概述SkillSub Pro - Python 题解与代码注释双功能技能是一项面向实际任务的技能,主要用于SkillSub Pro 是一个 Python 题解生成与代码注释的 双功能合体技能 ,专为学生、算法学习者和开发者设计;✅ 一个技能,两种用途 :;核心要点📝 题解模式 :输入题目/题号,自动生成完整 Python 题解(含详细注释、解题思路、复杂度分析);💬 注释模式 :输入 Python 代码,自动添加详细中。它将相关步骤、

下载

别忘了把结果从 GPU 移回 CPU 并转成 Python 值:

import torch
import torch.nn.functional as F

probs = F.softmax(logits, dim=-1)
pred_idx = torch.argmax(probs, dim=-1).item()
label_name = model.config.id2label[pred_idx]
confidence = probs[0][pred_idx].item()

注意:model.config.id2label 是模型自带的映射字典,不同 checkpoint 的 label 名称和顺序可能不同(例如 SST-2 是 {0: "NEGATIVE", 1: "POSITIVE"},而 IMDB 可能是 {0: "LABEL_0", 1: "LABEL_1"}),不能硬编码索引。

微调时用 Trainer 需显式传 compute_metrics,否则评估阶段无准确率输出

Trainer.train() 默认不计算任何指标,即使你传了 eval_dataset,日志里也只会显示 loss。想看 accuracy/f1,必须自己写 metric 函数并传入。

推荐用 datasets.load_metric("accuracy")(旧版)或 evaluate.load("accuracy")(新版):

from evaluate import load

metric = load("accuracy")

def compute_metrics(eval_pred):
    predictions, labels = eval_pred
    preds = np.argmax(predictions, axis=1)
    return metric.compute(predictions=preds, references=labels)

常见疏漏:

  • 忘记在 Trainer 初始化时传 compute_metrics=compute_metrics
  • 没对 predictionsargmax,直接喂给 metric(会报维度错误)
  • 用旧版 load_metric 但没装 datasets 包,或新版没装 evaluate

模型最后几层的分类头适配很脆弱——改 dataset 的 label2id 映射、漏掉 num_labels 参数、甚至 tokenizer 多加了个 add_special_tokens=True,都可能导致训练中途 loss 爆涨或评估全错。动手前先跑通一个最小验证样本,比调参更重要。

Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!

相关专题

更多
python打包成可执行文件
python打包成可执行文件

本专题为大家带来python打包成可执行文件相关的文章,大家可以免费的下载体验。

2023.07.20

1551

4

python能做什么
python能做什么

python能做的有:可用于开发基于控制台的应用程序、多媒体部分开发、用于开发基于Web的应用程序、使用python处理数据、系统编程等等。本专题为大家提供python相关的各种文章、以及下载和课程。

2023.07.25

3624

7

format在python中的用法
format在python中的用法

Python中的format是一种字符串格式化方法,用于将变量或值插入到字符串中的占位符位置。通过format方法,我们可以动态地构建字符串,使其包含不同值。php中文网给大家带来了相关的教程以及文章,欢迎大家前来阅读学习。

2023.07.31

1549

3

python教程
python教程

Python已成为一门网红语言,即使是在非编程开发者当中,也掀起了一股学习的热潮。本专题为大家带来python教程的相关文章,大家可以免费体验学习。

2023.08.03

20637

23

python环境变量的配置
python环境变量的配置

Python是一种流行的编程语言,被广泛用于软件开发、数据分析和科学计算等领域。在安装Python之后,我们需要配置环境变量,以便在任何位置都能够访问Python的可执行文件。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2023.08.04

2567

5

python eval
python eval

eval函数是Python中一个非常强大的函数,它可以将字符串作为Python代码进行执行,实现动态编程的效果。然而,由于其潜在的安全风险和性能问题,需要谨慎使用。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2023.08.04

2607

5

scratch和python区别
scratch和python区别

scratch和python的区别:1、scratch是一种专为初学者设计的图形化编程语言,python是一种文本编程语言;2、scratch使用的是基于积木的编程语法,python采用更加传统的文本编程语法等等。本专题为大家提供scratch和python相关的文章、下载、课程内容,供大家免费下载体验。

2023.08.11

1063

5

python合并两个列表
python合并两个列表

Python是一种强大的编程语言,具有许多方便的功能和工具。在Python中,有多种方法可以合并两个列表。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2023.08.10

576

4

python是前端还是后端
python是前端还是后端

Python属于前端也属于后端,其灵活性和丰富的生态系统使得开发人员能够在不同的领域中灵活运用。本专题为大家提供python相关的文章、下载、课程内容,供大家免费下载体验。

2023.08.11

2023

5

热门下载

更多
网站特效
/
网站源码
/
网站素材
/
前端模板

精品课程

更多
相关推荐
/
热门推荐
/
最新课程