大数跨境

文本分类:如何使用 BERT 构建文本分类模型?

文本分类:如何使用 BERT 构建文本分类模型? 知识代码AI
2026-09-12
5
导读:文本分类:如何使用 BERT 构建文本分类模型?一、课程背景本节课进入第二个 NLP 实战项目,使用 PyTorch + BERT 构建文本分类模型。

文本分类:如何使用 BERT 构建文本分类模型?


一、课程背景

本节课进入第二个 NLP 实战项目,使用 PyTorch + BERT 构建文本分类模型。BERT 是业界使用最广泛的深度学习 NLP 算法模型之一,掌握它之后学习其他基于 Attention 的模型(如 GPT)可以举一反三。[pdf_24]


二、文本分类问题背景

为什么要选 BERT?

文本分类有很多经典方法,但在复杂场景下需要 BERT 这样强大的工具:[pdf_24]

经典方法
特点
关键词统计
最简单粗暴
贝叶斯分类
基于条件概率,时至今日仍是好选择
支持向量机(SVM)
很长时间在 NLP 中占据统治地位
随机森林 / LDA / 神经网络
百家争鸣

新闻文本分类的三大挑战

挑战
说明
类别多
新闻种类一般 50 种甚至以上
数据不平衡
社会、经济、体育等类别文章多,少儿、医疗等类别少
多语言
新闻来源广泛,需支持多语言

三、BERT 原理与特点

定义

BERT(Bidirectional Encoder Representation from Transformers)即双向 Transformer 的 Encoder,是一种基于 Attention 方法的模型,在文本分类、自动对话、语义理解等十几项 NLP 任务上刷新了历史最好成绩。[pdf_24]

核心特点

特点
说明
多头注意力(Multi-Head Attention)
将模型分为多个头,形成多个子空间,关注不同方面的信息,捕捉更丰富的特征
MLM 训练方式
随机屏蔽部分 token,用其他未被屏蔽的 token 预测被屏蔽的 token
动态词向量
根据上下文不同,对同一 token 给出动态变化的向量(如"苹果"可指水果或品牌)
多语言优势
通过 WordPiece 覆盖上百种语言,无需了解各语言的分词规则

四、安装与模型准备

环境安装

pip install Transformers

使用 hugging face 的 PyTorch 版 Transformers 包。[pdf_24]

Transformers 关键文件

文件
作用
convert_BERT_original_tf2_checkpoint_to_PyTorch.py
将 TensorFlow 预训练模型转换为 PyTorch 模型
modeling_BERT.py
使用 BERT 的范例代码

预训练模型版本

模型
配置
参数
BERT-Base, Uncased
12-layer, 768-hidden, 12-heads
110M
BERT-Large, Uncased
24-layer, 1024-hidden, 16-heads
340M
BERT-Base, Multilingual Cased
12-layer, 768-hidden, 12-heads, 104 languages
110M

本课选择 BERT-Base, Multilingual Cased 版本,支持 104 种语言。[pdf_24]

转换后的三个文件

文件
作用
config.json
BERT 模型的配置文件,记录所有训练参数设置
pytorch_model.bin
模型文件本身
vocab.txt
词表文件,用于识别所支持语言的字符、字符串或单词

五、BERT 输入格式

BERT 输入需转化为三种向量:[pdf_24]

1. Token embeddings(词向量)

  • 第一个开头的 token 必须是 [CLS]
  • [CLS] 作为整篇文本的语义表示,用于文本分类等任务

2. Segment embeddings

  • 区分两句话(如问答任务中的问句和答句)
  • 在本课分类任务中只有一个句子

3. Position embeddings

  • 记录单词的位置信息

六、模型构建:BERTForSequenceClassification

网络结构

输入文本 → [BERT 模型] → pooled_output → [Dropout] → [全连接层] → 分类结果

核心代码

class BERTForSequenceClassification(BERTPreTrainedModel):
    def __init__(self, config):
        super().__init__(config)
        self.num_labels = config.num_labels          # 类别标签数量
        self.bert = BertModel(config)
        self.dropout = nn.Dropout(config.hidden_dropout_prob)  # 减少过拟合
        self.classifier = nn.Linear(config.hidden_size, config.num_labels)
        self.init_weights()

    def forward(self, input_ids, attention_mask=None, token_type_ids=None, ...):
        outputs = self.bert(input_ids, attention_mask=attention_mask, ...)
        pooled_output = outputs[1]                    # 经过 BERT 得到的中间输出
        pooled_output = self.dropout(pooled_output)   # 减少过拟合
        logits = self.classifier(pooled_output)       # 输出最后的分类结果

BERT 输出信息

输出
说明
last_hidden_state
最后一层隐藏层状态序列,shape=(batch_size, sequence_length, hidden_size),hidden_size=768
pooled_output
序列第一个 token([CLS])的最后一个隐藏层状态,shape=(batch_size, hidden_size)
hidden_states / attentions
其他可用信息

模型配置(config.json)

字段
说明
id2label
类别标签 → 类别名称的映射
label2id
类别名称 → 类别标签的映射
num_labels_cate
类别数量

七、数据准备

数据处理三要素

组件
作用
InputExample
记录单个训练数据的文本内容结构
DataProcessor
将训练数据集的文本表示为多个 InputExample 组成的数据集合
get_features
将 InputExample 数据转换成 BERT 能理解的数据结构

生成三种关键数据

# input_ids:记录输入 token 对应在 vocab.txt 中的 id 序号
input_ids = tokenizer.encode(example.text_a, add_special_tokens=True,
                              max_length=min(max_length, tokenizer.max_len))

# attention_mask:记录属于第一个句子的 token 信息
attention_mask = [1 if mask_padding_with_zero else 0] * len(input_ids)

# labels:记录文本类别的信息

八、模型训练

优化器:AdamW

BERT 使用 AdamW 优化器,需要对参数分组,对 bias 和 LayerNorm 不使用权重衰减:[pdf_24]

from transformers import AdamW

param_optimizer = list(model.named_parameters())
no_decay = ['bias''LayerNorm.bias''LayerNorm.weight']
optimizer_grouped_parameters = [
    {'params': [p for n, p in param_optimizer if not any(nd in n for nd in no_decay)]},
    {'params': [p for n, p in param_optimizer if any(nd in n for nd in no_decay)]}
]
optimizer = AdamW(optimizer_grouped_parameters, lr=args.learning_rate)

训练循环

for epoch in trange(0, args.num_train_epochs):
    model.train()        # 设置为训练状态
    for step, batch in enumerate(tqdm(train_dataLoader)):
        step_loss = training_step(batch)  # 训练核心环节
        tr_loss += step_loss[0]
        optimizer.step()
        optimizer.zero_grad()

训练核心环节:training_step

def training_step(batch):
    input_ids, token_type_ids, attention_mask, labels = batch
    input_ids = input_ids.to(device)
    token_type_ids = token_type_ids.to(device)
    attention_mask = attention_mask.to(device)
    labels = labels.to(device)

    logits = model(input_ids, token_type_ids=token_type_ids,
                   attention_mask=attention_mask, labels=labels)
    loss_fct = BCEWithLogitsLoss()
    loss = loss_fct(logits.view(-1, num_labels_cate),
                    labels.view(-1, num_labels_cate))
    loss.backward()
    return loss
关键点
说明
logits
通过网络得到的预测输出
loss
基于 logits 计算,用于梯度更新
loss 函数
**BCEWithLogitsLoss()**(二分类交叉熵 + Sigmoid)

九、思考题:长文本处理

问题:BERT 处理文本有最大长度要求(512),遇到长文本该怎么办?[pdf_24]

方法
说明
缺点
截断法
head截断(从开头)、tail截断(从结尾)、head+tail截断(各保留一部分)
较为暴力
Pooling法
对多个片段进行池化聚合
性能较差
压缩法
提取文本中有限的 segment
压缩效果有限

十、核心要点总结

知识点
要点
BERT 全称
Bidirectional Encoder Representation from Transformers
核心架构
双向 Transformer 的 Encoder,基于多头注意力机制
训练方式
MLM(Mask Language Model),随机屏蔽 token 后预测
动态词向量
同一 token 根据上下文动态变化
输入三向量
Token embeddings([CLS]开头)、Segment embeddings、Position embeddings
模型结构
BERT → pooled_output → Dropout → 全连接层 → 分类
关键输出
pooled_output = outputs[1],即[CLS]的隐藏状态
损失函数
BCEWithLogitsLoss
优化器
AdamW(参数分组,bias 和 LayerNorm 无权重衰减)
长文本处理
截断法、Pooling法、压缩法

十一、小结

  1. BERT 原理:基于 Transformer 的 Encoder,采用多头注意力和 MLM 训练方式,具有动态词向量和多语言支持的优势。
  2. 模型构建BERTForSequenceClassification 包含 BERT 模型 + Dropout + 全连接分类层。
  3. 数据准备:通过 InputExample、DataProcessor、get_features 将文本转换为 input_ids、attention_mask、token_type_ids。
  4. 训练过程:使用 AdamW 优化器 + BCEWithLogitsLoss 损失函数,训练循环包含两个 for 循环(epoch + batch)。
  5. 业务思考:新闻文本面临类别多、数据不平衡、多语言等问题,需结合数据预处理技巧。
  6. 作者建议:尽管 GitHub 上已有封装完善的 BERT 代码,仍建议好好看一下 Transformer 中的模型代码,对技术提升有非常大的助益。[pdf_24]

【声明】内容源于网络
0
0
知识代码AI
技术基底 机器视觉全栈 × 光学成像 × 图像处理算法 编程栈 C++/C#工业开发 | Python智能建模 工具链 Halcon/VisionPro工业部署 | PyTorch/TensorFlow模型炼金术 | 模型压缩&嵌入式移植
内容 407
粉丝 0
知识代码AI 技术基底 机器视觉全栈 × 光学成像 × 图像处理算法 编程栈 C++/C#工业开发 | Python智能建模 工具链 Halcon/VisionPro工业部署 | PyTorch/TensorFlow模型炼金术 | 模型压缩&嵌入式移植
总阅读7.1k
粉丝0
内容407