BERT模型架构初探解析:从Transformer到BERT-Classification
本文主要是学习[动手写 bert 系列] bert model architecture 模型架构初探(embedding + encoder + pooler)专栏第三个视频的知识归纳总结。深入理解BERT模型的每一个组件,从架构设计到参数量分析,全面掌握这一NLP里程碑模型的核心原理。
BERT(Bidirectional Encoder Representations from Transformers)自2018年问世以来,彻底改变了自然语言处理领域的格局。它提出的"预训练+微调"范式,使得NLP模型能够像CV领域使用ImageNet预训练模型一样,使用大规模的预训练语言模型作为基础。
1. BERT与Transformer关系
1.1 Transformer:Encoder-Decoder架构
# Transformer原始架构示意
Transformer(
encoder: 6层EncoderLayer, # 编码器部分
decoder: 6层DecoderLayer # 解码器部分
)

原始Transformer是为机器翻译等序列到序列(seq2seq)任务设计的,包含完整的编码器和解码器。编码器用于理解源语言文本,解码器用于生成目标语言文本。
1.2 BERT的创新:专注的编码器
BERT做出了一个关键的设计决策:只使用Transformer的编码器部分。
# BERT的架构选择
BERT = Transformer.encoder # 仅保留编码器

这一选择背后的逻辑非常深刻:
-
任务定位:BERT专注于文本理解任务,而非文本生成
-
双向注意力:编码器天然支持双向注意力,能看到完整上下文
-
预训练目标:MLM和NSP任务只需要编码器就能完成
1.3 核心差异对比
| 特性 | Transformer | BERT |
|---|---|---|
| 架构 | Encoder-Decoder | 仅Encoder |
| 注意力 | 编码器双向+解码器因果 | 完全双向 |
| 位置编码 | 正弦余弦公式 | 可学习嵌入 |
| 主要应用 | 机器翻译 | 文本理解 |
| 预训练 | 无预训练 | MLM + NSP |
2. BERT模型三大核心组件
2.1 Embeddings:文本的三维表示
BERT的嵌入层堪称工程与理论的完美结合,它将一个token转化为包含语义、位置、段落三重信息的向量表示。

# BERT嵌入层详细结构
BertEmbeddings(
# 1. 语义嵌入:词本身的意义
(word_embeddings): Embedding(30522, 768) # 23M参数
# 2. 位置嵌入:词在序列中的位置
(position_embeddings): Embedding(512, 768) # 39万参数
# 3. 段落嵌入:词属于哪个句子
(token_type_embeddings): Embedding(2, 768) # 1.5K参数
# 4. 归一化与正则化
(LayerNorm): LayerNorm((768,), eps=1e-12) # 稳定训练
(dropout): Dropout(p=0.1) # 防止过拟合
)
2.1.1 词嵌入:语义的基石
# 词嵌入矩阵示例
vocab_size = 30522 # BERT的词表大小
hidden_size = 768 # 隐藏维度
# 实际数据示例
word_embeddings = {
"[CLS]": [0.1, 0.2, ..., 0.768], # 分类token
"[SEP]": [0.3, 0.1, ..., 0.568], # 分隔token
"the": [0.4, 0.5, ..., -0.123], # 高频词
"##ing": [0.2, -0.1, ..., 0.456], # 子词片段
}
-
使用WordPiece分词:平衡词表大小和OOV问题
-
30522个token:包含完整单词和常见子词
-
##前缀:表示子词片段(如##ing)
2.1.2 位置嵌入:序列的秩序
与Transformer使用固定的正弦余弦函数不同,BERT采用可学习的位置嵌入:
# 位置嵌入矩阵示例
position_embeddings = {
0: [0.01, 0.02, ..., 0.768], # 第0个位置
1: [0.11, 0.12, ..., 0.668], # 第1个位置
...
511: [0.91, 0.82, ..., 0.168] # 第511个位置
}
为什么可学习更好?
-
灵活性:让模型自己学习最佳的位置表示
-
适应性:不同任务可能需要不同的位置编码模式
-
扩展性:理论上可以扩展到更长序列
2.1.3 段落嵌入:句子的边界
段落嵌入是BERT为了支持下一句预测(NSP)任务而设计的:
# 段落嵌入示例
token_type_embeddings = {
0: [0.1, 0.2, ..., 0.768], # 句子A的所有token
1: [0.3, 0.4, ..., 0.568], # 句子B的所有token
}
使用场景:
-
单句分类:所有token类型为0
-
句对任务:第一句token为0,第二句token为1
2.2 Encoder:12层Transformer的智慧堆叠
BERT-base的核心是12个完全相同的Transformer编码器层。每层都是一个复杂的特征提取器:
# 单层Transformer编码器结构
BertLayer(
# 第一部分:自注意力机制
(attention): BertAttention(
# 1.1 多头注意力计算
(self): BertSelfAttention(
(query): Linear(768→768) # Q矩阵
(key): Linear(768→768) # K矩阵
(value): Linear(768→768) # V矩阵
)
# 1.2 注意力输出处理
(output): BertSelfOutput(
(dense): Linear(768→768) # 线性变换
(LayerNorm): LayerNorm # 残差连接+归一化
)
)
# 第二部分:前馈神经网络
(intermediate): BertIntermediate(
(dense): Linear(768→3072) # 特征扩展
)
# 第三部分:前馈输出
(output): BertOutput(
(dense): Linear(3072→768) # 特征压缩
(LayerNorm): LayerNorm # 残差连接+归一化
)
)
2.2.1 多头自注意力机制详解
自注意力是Transformer的灵魂,让每个token都能"看到"序列中的所有其他token:
# 自注意力计算过程
def self_attention(x):
# 1. 计算Q、K、V
Q = linear(x, W_q) # [batch, seq_len, 768]
K = linear(x, W_k) # [batch, seq_len, 768]
V = linear(x, W_v) # [batch, seq_len, 768]
# 2. 拆分成12个头(每个头64维)
Q = reshape(Q, [batch, seq_len, 12, 64])
K = reshape(K, [batch, seq_len, 12, 64])
V = reshape(V, [batch, seq_len, 12, 64])
# 3. 每个头独立计算注意力
for head in range(12):
# 注意力分数 = softmax(Q·K^T / sqrt(64))
attention_scores = matmul(Q[head], transpose(K[head])) / 8.0
# 注意力权重
attention_weights = softmax(attention_scores)
# 加权求和
head_output = matmul(attention_weights, V[head])
# 4. 拼接12个头的输出
output = concat([head_0, head_1, ..., head_11])
return output
多头注意力的优势:
-
并行计算:12个头可以同时计算
-
多角度理解:每个头可能关注不同的语法或语义特征
-
表达能力:比单头注意力更强的表示能力
2.2.2 前馈网络:特征的深度加工
# 前馈网络计算
def feed_forward(x):
# 第一层:扩展维度(768→3072)
intermediate = gelu(linear(x, W_intermediate))
# 第二层:压缩维度(3072→768)
output = linear(intermediate, W_output)
return output
为什么需要前馈网络?
-
非线性变换:引入GELU激活函数
-
特征交互:让不同维度特征充分交互
-
容量提升:增加模型表达能力
2.2.3 残差连接与层归一化:训练的稳定器
这是BERT能够训练12层深度网络的关键:
# 残差连接 + 层归一化
def residual_layernorm(x, sublayer_output):
# 残差连接:保留原始信息
residual = x + sublayer_output
# 层归一化:稳定数值分布
output = layernorm(residual)
return output
作用:
-
缓解梯度消失:深层网络中梯度更容易传播
-
保留低层信息:避免高层特征覆盖底层特征
-
加速收敛:更稳定的梯度分布
2.2.4 12层的层级化理解
BERT的12层不是简单的重复,而是形成了层次化的语义理解:
# 12层Transformer的语义层次
layer_semantics = {
"layer_1-4": "表面特征层",
"功能": "捕获词性、基本语法",
"示例": "识别名词、动词、形容词"
"layer_5-8": "语义理解层",
"功能": "理解短语、简单语义",
"示例": "理解'红色苹果'、'快速奔跑'"
"layer_9-12": "推理抽象层",
"功能": "篇章理解、逻辑推理",
"示例": "理解文章主旨、推断作者意图"
}
2.3 Pooler:句子的指纹提取
虽然简单,但Pooler在句子级任务中起着关键作用:
BertPooler(
(dense): Linear(768→768) # 线性变换
(activation): Tanh() # 激活函数
)
工作原理:
def pooler_output(last_hidden_state):
# 1. 取[CLS] token的表示
cls_token = last_hidden_state[:, 0, :] # [batch_size, 768]
# 2. 线性变换 + Tanh激活
pooled = tanh(linear(cls_token)) # [batch_size, 768]
return pooled
为什么用[CLS] token?
-
预训练一致性:在NSP任务中,[CLS]就被训练为句子表示
-
位置固定:总是序列的第一个位置,便于提取
-
信息聚合:通过自注意力机制聚合了全句信息
3. BertModel与BertForSequenceClassification
3.1 BertModel:基础特征提取器
# BertModel的基本结构
BertModel(
(embeddings): BertEmbeddings # 嵌入层
(encoder): BertEncoder # 12层编码器
(pooler): BertPooler # 池化层
)
输出包含:
-
last_hidden_state:最后一层所有token的表示 -
pooler_output:[CLS] token的池化表示
3.2 BertForSequenceClassification:分类专用模型
# BertForSequenceClassification的完整结构
BertForSequenceClassification(
# 1. 共享的BERT基础模型
(bert): BertModel(
embeddings + encoder + pooler
)
# 2. 额外的Dropout层(增强正则化)
(dropout): Dropout(p=0.1)
# 3. 分类头(任务特定)
(classifier): Linear(in_features=768, out_features=2)
)
3.3 权重加载的智慧
当你加载预训练模型时,会看到这样的警告信息:
Some weights of the model checkpoint at bert-base-uncased were not used...
Some weights of BertForSequenceClassification were not initialized...
这实际上是BERT设计精妙之处:
3.3.1 丢弃预训练专用头
# 预训练时包含的权重(下游任务不需要)
预训练权重 = {
'cls.predictions.*': '掩码语言模型头',
'cls.seq_relationship.*': '下一句预测头'
}
# 下游任务保留的权重
基础模型权重 = {
'embeddings.*': '嵌入层权重',
'encoder.*': '编码器权重',
'pooler.*': '池化层权重'
}
3.3.2初始化任务特定头
# 新初始化的权重
新权重 = {
'classifier.weight': '随机初始化',
'classifier.bias': '零初始化'
}
# 原因:不同任务需要不同的分类头
# 文本分类:类别数不同
# 情感分析:二分类或多分类
# 命名实体识别:序列标注
4. BERT模型参数量解析(为什么是大模型?)
4.1 参数概览:总参数与可学习参数
当我们深入分析BERT-base模型时,首先需要明确两个关键概念:总参数和可学习参数。
参数统计核心代码:
# 从代码中提取的关键参数统计
total_params = 109,482,240 # 约1.09亿参数
total_learnable_params = 109,482,240 # 所有参数都是可学习的
# 打印参数信息的代码片段
for name, param in model.named_parameters():
print(name, '->', param.shape, '->', param.numel())
if param.requires_grad:
total_learnable_params += param.numel()
total_params += param.numel()
核心概念解析
1. 总参数 (Total Parameters)
-
定义:模型中所有权重和偏置的总数
-
BERT-base:109,482,240 个参数
-
意义:衡量模型复杂度和存储需求的主要指标
2. 可学习参数 (Learnable Parameters)
-
定义:在训练过程中通过反向传播更新的参数
-
BERT-base:109,482,240 个参数(全部可学习)
-
requires_grad=True:标记参数是否需要梯度更新
为什么所有参数都可学习?
-
预训练特性:BERT需要在大规模语料上进行预训练
-
端到端训练:从嵌入层到输出层都需要优化
-
迁移学习:微调时需要调整所有层适应下游任务
4.2 各组件参数详细分解
让我们按照模型的前向传播顺序,深入分析每个组件的参数构成:
4.2.1 Embeddings 层参数分析(占总参数21.77%)
# 嵌入层参数明细表
embeddings_parameters = {
'组件': {
'word_embeddings': {
'形状': '[30522, 768]',
'参数量': 23,440,896,
'占比': '21.41%',
'计算公式': '词表大小 × 隐藏维度'
},
'position_embeddings': {
'形状': '[512, 768]',
'参数量': 393,216,
'占比': '0.36%',
'计算公式': '最大序列长度 × 隐藏维度'
},
'token_type_embeddings': {
'形状': '[2, 768]',
'参数量': 1,536,
'占比': '0.0014%',
'计算公式': '句子类型数 × 隐藏维度'
},
'LayerNorm(嵌入层)': {
'形状': '权重[768], 偏置[768]',
'参数量': 1,536,
'占比': '0.0014%',
'计算公式': '隐藏维度 × 2'
}
},
'汇总': {
'总参数量': 23,837,184,
'占总参数比例': '21.77%',
'特点': '主要由词嵌入矩阵主导'
}
}
嵌入层参数特点:
-
词嵌入矩阵是最大的单个参数块
-
占嵌入层参数的98.34%
-
占模型总参数的21.41%
-
-
位置嵌入相对较小但关键
-
支持最大512个token的位置编码
-
可学习的比固定正弦余弦编码更灵活
-
-
段落嵌入极简设计
-
仅2×768=1,536个参数
-
但支持句子关系理解的重要功能
-
4.2.2 Encoder 层参数分析(占总参数77.62%)
BERT-base包含12个完全相同的Transformer编码器层,这是模型参数的主体部分。
单层Transformer参数构成
# 单层Transformer参数详细计算
single_layer_breakdown = {
'注意力机制(Attention)': {
'Q/K/V投影层': {
'每个': '768×768 = 589,824参数',
'三个合计': '1,769,472参数',
'偏置': '768×3 = 2,304参数'
},
'注意力输出层': {
'线性变换': '768×768 = 589,824参数',
'偏置': '768 = 768参数'
},
'LayerNorm(注意力)': {
'权重': '768参数',
'偏置': '768参数'
},
'注意力部分小计': '2,363,904参数'
},
'前馈网络(Feed-Forward)': {
'中间层扩展': {
'线性变换': '768×3072 = 2,359,296参数',
'偏置': '3072 = 3,072参数'
},
'输出层压缩': {
'线性变换': '3072×768 = 2,359,296参数',
'偏置': '768 = 768参数'
},
'LayerNorm(前馈)': {
'权重': '768参数',
'偏置': '768参数'
},
'前馈网络小计': '4,717,824参数'
},
'单层总计': {
'计算': '2,363,904 + 4,717,824 = 7,081,728参数',
'约等于': '708万参数'
}
}
12层Encoder总参数计算
# 12层参数汇总
encoder_total_params = {
'单层参数': 7,081,728,
'层数': 12,
'总计': '7,081,728 × 12 = 84,980,736参数',
'占总参数比例': '84,980,736 ÷ 109,482,240 = 77.62%',
'备注': '约8498万参数,是模型的主要部分'
}
编码器参数分布特点:
encoder_distribution = {
'注意力机制占比': {
'单层': '2,363,904 ÷ 7,081,728 = 33.38%',
'总12层': '28,366,848 ÷ 84,980,736 = 33.38%'
},
'前馈网络占比': {
'单层': '4,717,824 ÷ 7,081,728 = 66.62%',
'总12层': '56,613,888 ÷ 84,980,736 = 66.62%'
},
'关键发现': '前馈网络的参数是注意力机制的两倍'
}
4.2.3 Pooler 层参数分析(占总参数0.54%)
# 池化层参数分析
pooler_parameters = {
'结构': {
'dense层': {
'形状': '权重[768, 768],偏置[768]',
'参数量': '589,824 + 768 = 590,592',
'作用': '将[CLS]表示映射到句子向量'
}
},
'特点': {
'参数极少': '仅占模型总参数的0.54%',
'功能重要': '生成句子级表示用于分类任务',
'激活函数': 'Tanh,将输出限制在[-1, 1]范围'
}
}
5. 不同的BERT变体
# BERT家族成员比较
bert_family = {
'BERT-base': {
'参数': '110M',
'层数': 12,
'隐藏维': 768,
'头数': 12
},
'BERT-large': {
'参数': '340M',
'层数': 24,
'隐藏维': 1024,
'头数': 16
},
'DistilBERT': {
'参数': '66M',
'特点': '知识蒸馏,速度提升60%',
'应用': '移动端、实时推理'
},
'RoBERTa': {
'改进': '去除NSP,更多数据,更长训练',
'效果': '多项任务超越BERT'
}
}
6. 总结
经过这次初探解析,我们已经一起搞懂了BERT的核心奥秘:明白了BERT如何从Transformer的编码器演变而来,拆解了它的三大组件设计原理,清楚了基础模型与分类模型的差异,认识了1.09亿参数的分布规律,还了解了BERT家族的各种变体。如果你还有些地方觉得模糊,完全正常!学习就像训练神经网络,需要时间和迭代。让我们一起保持好奇,继续前行,在AI的星辰大海中,每个踏实的脚步都在引领我们走向更精彩的未来!加油,我们都可以成为更好的自己!💪
更多推荐



所有评论(0)