本文主要是学习[动手写 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?

  1. 预训练一致性:在NSP任务中,[CLS]就被训练为句子表示

  2. 位置固定:总是序列的第一个位置,便于提取

  3. 信息聚合:通过自注意力机制聚合了全句信息

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%',
        '特点': '主要由词嵌入矩阵主导'
    }
}

嵌入层参数特点:

  1. 词嵌入矩阵是最大的单个参数块

    • 占嵌入层参数的98.34%

    • 占模型总参数的21.41%

  2. 位置嵌入相对较小但关键

    • 支持最大512个token的位置编码

    • 可学习的比固定正弦余弦编码更灵活

  3. 段落嵌入极简设计

    • 仅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的星辰大海中,每个踏实的脚步都在引领我们走向更精彩的未来!加油,我们都可以成为更好的自己!💪

Logo

中国智能体开发者社区,聚焦智能体与大模型开发,提供前沿资讯、实用工具链、开源项目及行业案例。通过技术沙龙、开发者大赛等活动,促进经验交流与协作,助力开发者快速构建创新智能应用。

更多推荐