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

用文字拆解 BERT 的结构图
假设我们现在要画一个 BERT-base 的模型结构(12 层 Transformer),可以这样分层描述:
- 输入层
- 文本经过 WordPiece 分词器变成 token IDs(如
[CLS]我爱 NLP[SEP]→[101, 2769, 4263, 6639, 102]) -
每个 token 对应三个嵌入向量的和:
- Token Embeddings(词嵌入,维度 768)
- Position Embeddings(位置编码,最大长度 512)
- Segment Embeddings(句子分段,0/ 1 区分句子)
-
编码器堆叠(×12)
- 每层结构完全一致,包含:
- Multi-Head Attention(多头注意力,12 个 head)
- Layer Normalization(层标准化)
- Feed Forward Network(前馈网络,中间层维度 3072)
-
每层输入输出始终保持 768 维
-
输出层
- 取最后一层的
[CLS]位置向量作为句子表示 - 其他 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
基础使用三件套
-
加载模型与分词器
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) -
获取各层输出
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 -
可视化注意力权重
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
进阶探索建议
-
安装
bertviz观察注意力模式:from bertviz import model_view model_view(attention=outputs.attentions, tokens=tokenizer.convert_ids_to_tokens(inputs.input_ids[0])) -
比较 BERT 和 RoBERTa 的区别:
- RoBERTa 去掉了 NSP 任务
- 使用动态 masking 而非静态
- 更大的 batch size 和更多数据
理解 BERT 结构后,你会发现 Transformer 就像乐高积木——通过堆叠相同的编码器层,配合巧妙的注意力机制,最终实现强大的语言理解能力。建议从 HuggingFace 的 demo 开始,亲手修改参数观察输出变化,这种实践比死磕理论论文有效得多!
正文完
发表至: 人工智能
近一天内
