BERT模型微调实战:从零开始构建文本分类器

1次阅读
没有评论

共计 2461 个字符,预计需要花费 7 分钟才能阅读完成。

image.webp

为什么需要 BERT 微调?

在自然语言处理(NLP)任务中,预训练语言模型如 BERT 已经成为标配。但在实际业务场景中,我们往往需要针对特定任务进行微调。比如:

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
)

常见问题解决方案

类别不平衡

  1. 加权损失函数:

    weights = torch.tensor([1.0, 2.0, 1.5])  # 对少数类加大权重
    criterion = nn.CrossEntropyLoss(weight=weights)

  2. 过采样少数类

  3. 欠采样多数类

GPU 显存不足

  1. 减小 batch size
  2. 使用梯度累积
  3. 启用混合精度
  4. 尝试模型蒸馏
  5. 使用更小的 BERT 变体(如 DistilBERT)

识别过拟合

  • 训练 loss 持续下降但验证 loss 不降或上升
  • 早停法 (early stopping) 是最直接的对策

延伸思考

  1. 领域自适应策略:可以在目标领域数据上继续预训练(继续 MLM 任务),再进行微调
  2. LoRA 等高效微调技术:
  3. 优点:大幅减少可训练参数
  4. 缺点:可能需要更多调参

通过以上步骤,你应该能够成功构建一个 BERT 文本分类器。实践中遇到问题时,不妨回到这些基础方法进行调整。

正文完
 0
评论(没有评论)