BERT预训练模型结构图解析:从零理解Transformer核心架构

1次阅读
没有评论

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

image.webp

为什么初学者觉得 BERT 难懂?

刚开始接触 BERT 时,我被一堆术语搞得头晕——self-attention(自注意力)、positional encoding(位置编码)、encoder layers(编码器层)……明明论文里的结构图看起来挺整齐,但具体数据怎么流动的?各层之间如何交互?这些细节让我这种菜鸟非常困惑。后来发现,其实只要抓住几个关键模块,理解起来就会轻松很多。

BERT 预训练模型结构图解析:从零理解 Transformer 核心架构

用文字拆解 BERT 的结构图

假设我们现在要画一个 BERT-base 的模型结构(12 层 Transformer),可以这样分层描述:

  1. 输入层
  2. 文本经过 WordPiece 分词器变成 token IDs(如 [CLS] 我爱 NLP[SEP][101, 2769, 4263, 6639, 102]
  3. 每个 token 对应三个嵌入向量的和:

    • Token Embeddings(词嵌入,维度 768)
    • Position Embeddings(位置编码,最大长度 512)
    • Segment Embeddings(句子分段,0/ 1 区分句子)
  4. 编码器堆叠(×12)

  5. 每层结构完全一致,包含:
    • Multi-Head Attention(多头注意力,12 个 head)
    • Layer Normalization(层标准化)
    • Feed Forward Network(前馈网络,中间层维度 3072)
  6. 每层输入输出始终保持 768 维

  7. 输出层

  8. 取最后一层的 [CLS] 位置向量作为句子表示
  9. 其他 token 位置输出可用于序列标注任务

两个预训练任务揭秘

BERT 通过两种玩法预训练:

  • Masked Language Model (MLM)
  • 随机盖住 15% 的 token(其中 80% 替换为[MASK],10% 随机替换,10% 保持不变)
  • 让模型预测被盖住的词,比如:

    输入:"人工 [MASK] 是 AI 的重要组成部分"
    目标:预测 "智能"

  • Next Sentence Prediction (NSP)

  • 给两个句子,判断是否连续:
    • 正样本:” 今天天气真好 ” + “ 适合去公园散步 ”
    • 负样本:” 深度学习很强大 ” + “ 西红柿炒鸡蛋的做法 ”
  • [CLS] 位置的输出做二分类

动手实践:用 HuggingFace 玩转 BERT

先安装必要库:

pip install transformers torch matplotlib

基础使用三件套

  1. 加载模型与分词器

    from transformers import BertTokenizer, BertModel
    import torch
    
    # 初始化
    model_name = "bert-base-uncased"
    tokenizer = BertTokenizer.from_pretrained(model_name)
    model = BertModel.from_pretrained(model_name)
    
    # 编码文本
    text = "I love natural language processing"
    inputs = tokenizer(text, return_tensors="pt")  # 返回 PyTorch 张量
    print(inputs.input_ids.shape)  # 输出:torch.Size([1, 6]) (batch_size, seq_len)

  2. 获取各层输出

    with torch.no_grad():
        outputs = model(**inputs, output_hidden_states=True)
    
    # 所有隐藏层(13 层,包含初始嵌入层)all_hidden_states = outputs.hidden_states  # tuple 形状 (13, batch_size, seq_len, hidden_dim)
    
    # 最后一层的 CLS 向量
    last_hidden_state = outputs.last_hidden_state  # (1, 6, 768)
    cls_embedding = last_hidden_state[0, 0, :]  # 取第一个样本的第一个 token

  3. 可视化注意力权重

    import matplotlib.pyplot as plt
    
    # 获取第 3 层第 1 个 head 的注意力权重
    attention = outputs.attentions[2][0, 0].numpy()  # shape (seq_len, seq_len)
    
    # 画热力图
    plt.imshow(attention, cmap='hot')
    plt.colorbar()
    plt.xlabel("Key Position")
    plt.ylabel("Query Position")
    plt.title("Attention Heatmap (Layer 3 Head 1)")
    plt.show()

新手避雷指南

  • 长文本截断问题
  • BERT 默认最大长度 512,超长文本需要截断
  • 解决方案:

    # 手动截断
    inputs = tokenizer(text, max_length=512, truncation=True)

  • [CLS]的正确用法

  • 分类任务用pooler_output(已通过 tanh 激活)
  • 不要直接拿最后一层的 [CLS] 向量当句子表示
    # 正确获取句子表示
    sentence_embedding = outputs.pooler_output

进阶探索建议

  1. 安装 bertviz 观察注意力模式:

    from bertviz import model_view
    model_view(attention=outputs.attentions, tokens=tokenizer.convert_ids_to_tokens(inputs.input_ids[0]))

  2. 比较 BERT 和 RoBERTa 的区别:

  3. RoBERTa 去掉了 NSP 任务
  4. 使用动态 masking 而非静态
  5. 更大的 batch size 和更多数据

理解 BERT 结构后,你会发现 Transformer 就像乐高积木——通过堆叠相同的编码器层,配合巧妙的注意力机制,最终实现强大的语言理解能力。建议从 HuggingFace 的 demo 开始,亲手修改参数观察输出变化,这种实践比死磕理论论文有效得多!

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