文本分类:如何使用 BERT 构建文本分类模型?
一、课程背景
本节课进入第二个 NLP 实战项目,使用 PyTorch + BERT 构建文本分类模型。BERT 是业界使用最广泛的深度学习 NLP 算法模型之一,掌握它之后学习其他基于 Attention 的模型(如 GPT)可以举一反三。[pdf_24]
二、文本分类问题背景
为什么要选 BERT?
文本分类有很多经典方法,但在复杂场景下需要 BERT 这样强大的工具:[pdf_24]
新闻文本分类的三大挑战
三、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 |
|
预训练模型版本
|
|
|
|
|
|
12-layer, 768-hidden, 12-heads
|
|
|
|
24-layer, 1024-hidden, 16-heads
|
|
| BERT-Base, Multilingual Cased |
12-layer, 768-hidden, 12-heads, 104 languages
|
|
本课选择 BERT-Base, Multilingual Cased 版本,支持 104 种语言。[pdf_24]
转换后的三个文件
|
|
|
| config.json |
|
| pytorch_model.bin |
|
| vocab.txt |
|
五、BERT 输入格式
BERT 输入需转化为三种向量:[pdf_24]
1. Token embeddings(词向量)
-
-
[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 |
|
| loss 函数 |
**BCEWithLogitsLoss()**(二分类交叉熵 + Sigmoid)
|
九、思考题:长文本处理
问题:BERT 处理文本有最大长度要求(512),遇到长文本该怎么办?[pdf_24]
|
|
|
|
| 截断法 |
head截断(从开头)、tail截断(从结尾)、head+tail截断(各保留一部分)
|
|
| Pooling法 |
|
|
| 压缩法 |
|
|
十、核心要点总结
|
|
|
| BERT 全称 |
Bidirectional Encoder Representation from Transformers
|
| 核心架构 |
双向 Transformer 的 Encoder,基于多头注意力机制
|
| 训练方式 |
MLM(Mask Language Model),随机屏蔽 token 后预测
|
| 动态词向量 |
|
| 输入三向量 |
Token embeddings([CLS]开头)、Segment embeddings、Position embeddings
|
| 模型结构 |
BERT → pooled_output → Dropout → 全连接层 → 分类
|
| 关键输出 |
pooled_output = outputs[1],即[CLS]的隐藏状态
|
| 损失函数 |
|
| 优化器 |
AdamW(参数分组,bias 和 LayerNorm 无权重衰减)
|
| 长文本处理 |
|
十一、小结
-
BERT 原理:基于 Transformer 的 Encoder,采用多头注意力和 MLM 训练方式,具有动态词向量和多语言支持的优势。
-
模型构建:
BERTForSequenceClassification 包含 BERT 模型 + Dropout + 全连接分类层。
-
数据准备:通过 InputExample、DataProcessor、get_features 将文本转换为 input_ids、attention_mask、token_type_ids。
-
训练过程:使用 AdamW 优化器 + BCEWithLogitsLoss 损失函数,训练循环包含两个 for 循环(epoch + batch)。
-
业务思考:新闻文本面临类别多、数据不平衡、多语言等问题,需结合数据预处理技巧。
-
作者建议:尽管 GitHub 上已有封装完善的 BERT 代码,仍建议好好看一下 Transformer 中的模型代码,对技术提升有非常大的助益。[pdf_24]