共计 2461 个字符,预计需要花费 7 分钟才能阅读完成。
为什么需要 BERT 微调?
在自然语言处理(NLP)任务中,预训练语言模型如 BERT 已经成为标配。但在实际业务场景中,我们往往需要针对特定任务进行微调。比如:

- 客服工单分类:将用户反馈自动分类为 ” 技术问题 ”、” 账单问题 ”、” 账户问题 ” 等,大幅提高客服效率
- 新闻主题识别:自动标注新闻属于 ” 政治 ”、” 经济 ”、” 体育 ” 等类别,便于内容管理和推荐
这些任务都需要模型理解特定领域的语义,这正是 BERT 微调的价值所在。
微调 vs 特征提取
| 方法 | 原理 | 适用场景 | 计算成本 |
|---|---|---|---|
| Fine-tuning | 调整所有模型参数 | 数据量较大(>10k 样本) | 高 |
| Feature-based | 固定 BERT 参数,仅训练分类层 | 数据量小(<1k 样本) | 低 |
核心实现流程
1. 环境准备
首先安装必要的库:
pip install transformers==4.28.1 torch==1.13.1 pytorch-lightning==1.9.0
2. 加载预训练模型
from transformers import BertTokenizer, BertForSequenceClassification
# 加载分词器和模型
tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')
model = BertForSequenceClassification.from_pretrained(
'bert-base-uncased',
num_labels=5, # 分类类别数
output_attentions=False,
output_hidden_states=False
)
3. 数据预处理
from torch.utils.data import Dataset
class TextDataset(Dataset):
def __init__(self, texts: list[str], labels: list[int], tokenizer, max_len: int = 512):
self.texts = texts
self.labels = labels
self.tokenizer = tokenizer
self.max_len = max_len
def __len__(self):
return len(self.texts)
def __getitem__(self, idx) -> dict:
text = str(self.texts[idx])
label = self.labels[idx]
# 关键:处理文本截断和特殊 token
encoding = self.tokenizer.encode_plus(
text,
add_special_tokens=True,
max_length=self.max_len,
truncation=True,
padding='max_length',
return_attention_mask=True,
return_tensors='pt'
)
return {'input_ids': encoding['input_ids'].flatten(),
'attention_mask': encoding['attention_mask'].flatten(),
'labels': torch.tensor(label, dtype=torch.long)
}
4. 自定义 DataCollator
from transformers import DataCollatorWithPadding
# 继承并扩展默认的 DataCollator
class CustomDataCollator(DataCollatorWithPadding):
def __call__(self, features):
batch = super().__call__(features)
# 确保 labels 存在且格式正确
if 'labels' in features[0]:
batch['labels'] = torch.tensor([f['labels'] for f in features])
return batch
性能优化技巧
混合精度训练
from pytorch_lightning import Trainer
# 在 Trainer 中启用混合精度
trainer = Trainer(
precision=16, # 使用 fp16
accelerator='gpu',
devices=1
)
梯度累积
trainer = Trainer(
accumulate_grad_batches=4, # 每 4 个 batch 更新一次梯度
# ... 其他参数
)
学习率 warmup
数学原理:线性或余弦式逐步提高学习率,避免初期大梯度破坏预训练权重。
from transformers import get_linear_schedule_with_warmup
# 在训练循环中
optimizer = AdamW(model.parameters(), lr=2e-5)
scheduler = get_linear_schedule_with_warmup(
optimizer,
num_warmup_steps=100, # warmup 步数
num_training_steps=total_steps
)
常见问题解决方案
类别不平衡
-
加权损失函数:
weights = torch.tensor([1.0, 2.0, 1.5]) # 对少数类加大权重 criterion = nn.CrossEntropyLoss(weight=weights) -
过采样少数类
- 欠采样多数类
GPU 显存不足
- 减小 batch size
- 使用梯度累积
- 启用混合精度
- 尝试模型蒸馏
- 使用更小的 BERT 变体(如 DistilBERT)
识别过拟合
- 训练 loss 持续下降但验证 loss 不降或上升
- 早停法 (early stopping) 是最直接的对策
延伸思考
- 领域自适应策略:可以在目标领域数据上继续预训练(继续 MLM 任务),再进行微调
- LoRA 等高效微调技术:
- 优点:大幅减少可训练参数
- 缺点:可能需要更多调参
通过以上步骤,你应该能够成功构建一个 BERT 文本分类器。实践中遇到问题时,不妨回到这些基础方法进行调整。
正文完
