从零构建Transformer Encoder:AG News文本分类任务实战指南

1次阅读
没有评论

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

image.webp

背景与痛点

刚接触 Transformer 架构时,很多同学会被其复杂的结构吓退。特别是实现 Encoder 部分时,最常见的三大拦路虎是:

从零构建 Transformer Encoder:AG News 文本分类任务实战指南

  • 位置编码的理解:为什么不能直接用传统 RNN 的序列处理方式?正弦 / 余弦位置编码的数学意义是什么?

  • 注意力机制实现:QKV 矩阵究竟如何计算?为什么需要缩放点积注意力?

  • 训练不稳定问题:模型初期 loss 震荡剧烈,甚至出现 NaN 值该怎么办?

这些痛点在实际做 AG News 分类任务时会被放大——我们需要处理可变长度的新闻文本,同时要保证模型对关键词的捕捉能力。

技术选型:库 vs 原生实现

面对两种主流方案,我的建议是:

  1. HuggingFace 优势
  2. 三行代码调用预训练模型
  3. 内置优化过的 Attention 计算
  4. 适合快速原型验证

  5. 原生 PyTorch 优势

  6. 彻底掌握模型细节
  7. 方便自定义修改(比如调整 Encoder 层数)
  8. 更轻量无依赖

考虑到本文的教学目的,我们选择从零实现。放心,我会带你避开所有深坑!

核心实现四步走

第一步:数据预处理

AG News 数据集包含 4 类新闻标题和描述,我们需要:

from torchtext.datasets import AG_NEWS
from torchtext.data.utils import get_tokenizer

tokenizer = get_tokenizer('basic_english')
train_iter = AG_NEWS(split='train')

# 构建词汇表
vocab = build_vocab_from_iterator(map(tokenizer, [text for label, text in train_iter]),
    specials=['<unk>', '<pad>']
)
vocab.set_default_index(vocab['<unk>'])

# 文本向量化函数
def text_pipeline(text):
    return vocab(tokenizer(text))

关键点说明:

  • 使用基础英文分词器处理标点
  • 预留 <unk><pad>两个特殊 token
  • 最终生成形如 [23, 156, 792] 的数值序列

第二步:实现 Encoder 层

精简版 SingleHeadAttention 实现(完整版见后续代码):

import torch
import torch.nn as nn

class SelfAttention(nn.Module):
    def __init__(self, embed_size):
        super().__init__()
        self.query = nn.Linear(embed_size, embed_size)
        self.key = nn.Linear(embed_size, embed_size)
        self.value = nn.Linear(embed_size, embed_size)
        self.softmax = nn.Softmax(dim=-1)

    def forward(self, x):
        Q = self.query(x)
        K = self.key(x)
        V = self.value(x)

        # 缩放点积注意力
        scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(Q.size(-1)))
        attention = self.softmax(scores)
        return torch.matmul(attention, V)

第三步:组装完整模型

class TransformerClassifier(nn.Module):
    def __init__(self, vocab_size, embed_size=128, num_classes=4):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, embed_size)
        self.position_encoding = PositionalEncoding(embed_size)  # 需自行实现
        self.encoder_layer = nn.TransformerEncoderLayer(
            d_model=embed_size, 
            nhead=8,
            dim_feedforward=512
        )
        self.transformer_encoder = nn.TransformerEncoder(self.encoder_layer, num_layers=3)
        self.fc = nn.Linear(embed_size, num_classes)

    def forward(self, x):
        x = self.embedding(x)
        x = self.position_encoding(x)
        x = self.transformer_encoder(x)
        x = x.mean(dim=1)  # 全局平均池化
        return self.fc(x)

第四步:训练技巧

三个关键超参数设置:

  1. 学习率:使用带 warmup 的 AdamW 优化器

    optimizer = AdamW(model.parameters(), lr=5e-5)
    scheduler = get_linear_schedule_with_warmup(optimizer, num_warmup_steps=100, num_training_steps=1000)

  2. Batch Size:根据 GPU 显存选择(建议 32-64)

  3. 梯度裁剪:防止梯度爆炸

    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

五大避坑指南

  1. 梯度消失问题
  2. 解决方案:每层添加残差连接
  3. 代码实现:

    x = x + self.attention(x)  # 残差连接

  4. 过拟合现象

  5. 对策:在 Embedding 后添加 Dropout 层
  6. 推荐参数:nn.Dropout(p=0.1)

  7. 位置编码失效

  8. 关键检查:确保 PE 值范围与 Embedding 匹配
  9. 调试方法:可视化前几个位置的编码向量

  10. 内存溢出(OOM)

  11. 应急方案:

    torch.cuda.empty_cache()
    reduce_batch_size()

  12. 预测时结果随机

  13. 根本原因:忘记 model.eval() 模式
  14. 完整预测流程:
    with torch.no_grad():
        model.eval()
        output = model(input)

进阶实战建议

当模型准确率稳定在 90%+ 后,可以考虑:

  • 模型压缩:使用知识蒸馏技术,将大模型的能力迁移到小模型

    student_model = TinyTransformer()
    distil_loss = KLDivLoss(teacher_logits, student_logits)

  • 部署优化:转换为 ONNX 格式提升推理速度

    torch.onnx.export(model, dummy_input, "ag_news.onnx")

思考题

  1. 如果新闻文本特别长(如超过 512 个 token),应该如何修改当前架构?
  2. 如何修改注意力机制使其能捕捉局部特征(类似 CNN 的效果)?
  3. 在多语言场景下,位置编码需要做哪些特殊处理?

希望这篇笔记能帮你打通 Transformer 的任督二脉!遇到问题欢迎在评论区交流~

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