共计 2742 个字符,预计需要花费 7 分钟才能阅读完成。
为什么需要 Transformer 做文本分类?
传统 NLP 方法如 TF-IDF+SVM 在文本分类中会遇到几个典型问题:

- 难以捕捉长距离语义依赖(超过 10 个词就效果下降)
- 无法理解一词多义(比如 ” 苹果 ” 在手机和水果场景中的区别)
- 需要人工设计特征工程(如 n -gram 组合)
而 Transformer 通过 Self-Attention 机制实现了:
- 任意距离的语义关联(无论词间距多远都能建立联系)
- 动态上下文理解(同一个词在不同位置有不同向量表示)
- 端到端特征学习(自动提取分层语义特征)
CLS Token 的魔法作用
在 BERT 类模型中,CLS(Classification)Token 是个特殊存在:
- 位于序列开头:[CLS] 我 爱 自然 语言 处理 [SEP]
- 训练时会吸收整个序列的语义信息
-
相比 Mean-Pooling 有三大优势:
-
避免有效信息被无关词稀释(比如停用词)
- 保留完整的分类决策信号(专门为分类任务优化)
- 支持跨序列交互(在 NLI 等任务中表现更好)
实验显示在 IMDB 影评数据集上:
| 方法 | 准确率 | 训练速度 |
|---|---|---|
| Mean-Pooling | 91.2% | 1.3x |
| CLS Token | 92.7% | 1.0x |
动手实现完整流程
数据预处理
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained('bert-base-uncased')
# 典型文本处理流程
def preprocess(texts, labels, max_len=128):
inputs = tokenizer(
texts,
padding='max_length',
truncation=True,
max_length=max_len,
return_tensors='pt'
)
# 注意 CLS Token 会自动插入到开头
return {'input_ids': inputs['input_ids'],
'attention_mask': inputs['attention_mask'],
'labels': torch.tensor(labels)
}
模型定义关键点
import torch.nn as nn
from transformers import BertModel
class BertClassifier(nn.Module):
def __init__(self, num_classes):
super().__init__()
self.bert = BertModel.from_pretrained('bert-base-uncased')
# 关键:仅用 CLS 位置对应的输出
self.classifier = nn.Linear(768, num_classes)
def forward(self, input_ids, attention_mask):
outputs = self.bert(
input_ids=input_ids,
attention_mask=attention_mask
)
# 取 CLS 位置的 hidden_state
cls_output = outputs.last_hidden_state[:, 0, :]
return self.classifier(cls_output)
训练技巧三件套
-
动态学习率:
from transformers import get_linear_schedule_with_warmup scheduler = get_linear_schedule_with_warmup( optimizer, num_warmup_steps=100, num_training_steps=total_steps ) -
早停机制:
if val_loss < best_loss: best_loss = val_loss patience = 0 torch.save(model.state_dict(), 'best_model.pt') else: patience += 1 if patience >= 3: break -
混合精度训练(省显存):
scaler = torch.cuda.amp.GradScaler() with torch.cuda.amp.autocast(): outputs = model(**batch) loss = criterion(outputs, batch['labels']) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()
生产环境调优经验
内存与效果平衡
- 长文本处理策略:
- 优先截断尾部(通常比头部信息密度低)
-
尝试滑动窗口 +Pooling(适合合同等长文档)
-
Batch Size 黄金法则:
# GPU 显存与 batch_size 关系估算 max_len = 256 # 序列长度 batch_size = (GPU_MEMORY - 1500) // (max_len * 2.5)
类别不平衡解决方案
-
损失函数加权:
weight = torch.tensor([1.0, 5.0]) # 少数类权重提升 criterion = nn.CrossEntropyLoss(weight=weight) -
过采样 + 欠采样组合:
from imblearn.over_sampling import RandomOverSampler ros = RandomOverSampler() X_resampled, y_resampled = ros.fit_resample(features.cpu().numpy(), labels.cpu().numpy() )
进阶应用方向
多标签分类改造
只需修改两点:
-
输出层替换为多个 Sigmoid:
self.classifier = nn.Linear(768, num_classes) # 计算 loss 时用 BCEWithLogitsLoss -
训练时用阈值判定:
threshold = 0.3 probs = torch.sigmoid(outputs) predictions = (probs > threshold).long()
序列标注任务适配
放弃 CLS Token,改用全部输出:
# 修改模型输出
outputs = self.bert(...).last_hidden_state # [batch, seq_len, dim]
logits = self.classifier(outputs) # 每个位置预测标签
实践建议
- 小数据场景优先微调预训练模型
- 使用 HuggingFace 的 Trainer 简化流程
- 监控 CLS Token 的 attention 分布(常出现异常聚焦时需检查)
完整可运行代码已放在:
Colab 实践链接
推荐延伸阅读:
–《Attention Is All You Need》原论文
– HuggingFace 官方 Fine-tuning 教程
– BERT 可视化工具 BertViz
正文完
发表至: 人工智能
近一天内
