CLS Transformer 入门指南:从零构建你的第一个文本分类模型

1次阅读
没有评论

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

image.webp

为什么需要 Transformer 做文本分类?

传统 NLP 方法如 TF-IDF+SVM 在文本分类中会遇到几个典型问题:

CLS Transformer 入门指南:从零构建你的第一个文本分类模型

  • 难以捕捉长距离语义依赖(超过 10 个词就效果下降)
  • 无法理解一词多义(比如 ” 苹果 ” 在手机和水果场景中的区别)
  • 需要人工设计特征工程(如 n -gram 组合)

而 Transformer 通过 Self-Attention 机制实现了:

  1. 任意距离的语义关联(无论词间距多远都能建立联系)
  2. 动态上下文理解(同一个词在不同位置有不同向量表示)
  3. 端到端特征学习(自动提取分层语义特征)

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)

训练技巧三件套

  1. 动态学习率:

    from transformers import get_linear_schedule_with_warmup
    
    scheduler = get_linear_schedule_with_warmup(
        optimizer,
        num_warmup_steps=100,
        num_training_steps=total_steps
    )

  2. 早停机制:

    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

  3. 混合精度训练(省显存):

    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) 

类别不平衡解决方案

  1. 损失函数加权:

    weight = torch.tensor([1.0, 5.0])  # 少数类权重提升
    criterion = nn.CrossEntropyLoss(weight=weight)

  2. 过采样 + 欠采样组合:

    from imblearn.over_sampling import RandomOverSampler
    
    ros = RandomOverSampler()
    X_resampled, y_resampled = ros.fit_resample(features.cpu().numpy(), 
        labels.cpu().numpy()
    )

进阶应用方向

多标签分类改造

只需修改两点:

  1. 输出层替换为多个 Sigmoid:

    self.classifier = nn.Linear(768, num_classes)
    # 计算 loss 时用 BCEWithLogitsLoss

  2. 训练时用阈值判定:

    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)  # 每个位置预测标签 

实践建议

  1. 小数据场景优先微调预训练模型
  2. 使用 HuggingFace 的 Trainer 简化流程
  3. 监控 CLS Token 的 attention 分布(常出现异常聚焦时需检查)

完整可运行代码已放在:
Colab 实践链接

推荐延伸阅读:
–《Attention Is All You Need》原论文
– HuggingFace 官方 Fine-tuning 教程
– BERT 可视化工具 BertViz

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