001-Transformer架构与Self-Attention机制深度解析

第1章 LLM基础原理

第1章第1节:Transformer架构与Self-Attention机制深度解析

本节目录

  1. 1.1 引言:从序列模型到注意力机制
  2. 1.2 Self-Attention原理
  3. 1.3 QKV矩阵计算详解
  4. 1.4 Scaled Dot-Product Attention
  5. 1.5 Multi-Head Attention机制
  6. 1.6 位置编码(Positional Encoding)
  7. 1.7 完整Transformer架构
  8. 1.8 前馈网络与残差连接
  9. 1.9 Python代码实现
  10. 1.10 架构对比
  11. 1.11 关键要点总结

1.1 引言:从序列模型到注意力机制

2017年,Google Brain团队在论文 "Attention Is All You Need" (Vaswani et al.)中提出了 Transformer 架构,彻底改变了自然语言处理领域的格局。 在此之前,序列建模任务(机器翻译、文本摘要、语言模型)几乎都依赖 **循环神经网络(RNN)**或 长短期记忆网络(LSTM)

RNN系列模型存在两个根本性瓶颈:

  • 顺序依赖性(Sequential Dependency):RNN必须按时间步逐步处理序列, 无法并行计算,导致训练速度极慢;
  • 长距离依赖问题(Long-Range Dependency):即便LSTM引入了门控机制, 对于超过数百个token的长序列,梯度仍会消失或爆炸,难以捕捉远距离的语义关联。

历史背景

注意力机制(Attention)并非Transformer首创。2014年,Bahdanau等人在神经机器翻译中 提出了软注意力(Soft Attention),允许解码器在每个时间步"关注"编码器的不同位置。 Transformer的革命性在于:完全抛弃循环结构,仅用注意力机制本身来捕捉所有 的序列关系。

Transformer的核心洞察是:序列中任意两个位置之间的依赖关系, 可以通过Self-Attention(自注意力)在单次前向传播中以 *O(1)*步骤直接计算------而RNN需要 *O(n)*步骤才能将信息从位置1传递到位置n。

这一突破带来的影响是深远的。从BERT、GPT系列到T5、PaLM、LLaMA, 几乎所有现代大型语言模型(LLM)都建立在Transformer架构之上。 理解Transformer的内部机制,是掌握现代AI工程的第一步。

1.2 Self-Attention原理

1.2.1 直觉理解

Self-Attention (自注意力)的核心思想可以用一个问题来概括: "在处理当前词时,序列中的哪些其他词对理解它最为重要?"

考虑句子:"The animal didn't cross the street because it was too tired."

当我们处理代词 "it" 时,Self-Attention允许模型自动"看向" "animal",赋予其更高的注意力权重,从而正确解析指代关系。 这种能力在RNN中依赖隐藏状态的线性传递,在Transformer中则通过直接的 点积相似度计算来实现。

1.2.2 数学基础:相似度度量

Self-Attention的本质是加权聚合(Weighted Aggregation): 给定一组值(Values),根据查询(Query)与每个键(Key)的相似度, 对这些值进行加权求和。

设序列长度为 n ,每个位置的表示维度为 d 。对于位置 i 处的token:

直觉公式(非标准化版本):

ini 复制代码
  output_i = Σ_j  similarity(i, j) × value_j

其中 similarity(i, j) 度量位置 i 的查询与位置 j 的键之间的相关性。

1.2.3 从加性注意力到点积注意力

计算相似度的方式主要有两种:

  • 加性注意力(Additive Attention / Bahdanau Attention)score(q, k) = v^T · tanh(W_q · q + W_k · k), 通过一个单层前馈网络计算,参数量较多;
  • 点积注意力(Dot-Product Attention)score(q, k) = q^T · k, 计算简单,可用高度优化的矩阵乘法实现。

Transformer选择了点积注意力,并加入了缩放因子 1/√d_k, 得到Scaled Dot-Product Attention。 缩放的原因将在下一节详述。

工程直觉

点积注意力在GPU上可以完全表示为矩阵乘法(GEMM),这使得现代深度学习框架 (PyTorch、JAX)能够利用高度优化的BLAS库和专用硬件加速(Tensor Core)实现极致性能。 这是Transformer在工程上能够扩展到数千亿参数的重要原因之一。

1.2.4 注意力矩阵的结构

对于长度为 n 的序列,Self-Attention会计算一个 n × n注意力矩阵 (Attention Matrix), 其中第 (i, j) 个元素表示位置 i 对位置 j 的注意力权重。

这个矩阵具有以下特性:

  • 每一行的权重经过Softmax归一化,行和为1;
  • 权重值越大,表示对应位置的信息对当前位置越重要;
  • 在Decoder的Masked Self-Attention中,上三角部分被遮蔽(设为-∞), 确保位置 i 只能看到位置 1..i(自回归约束)。

计算复杂度警告

Standard Self-Attention的时间复杂度和空间复杂度均为 O(n²·d) ,其中 n 为序列长度,d 为模型维度。 当序列长度超过4096时,这一二次方复杂度会成为显著瓶颈。 这正是Longformer、BigBird、FlashAttention等研究的出发点。 在设计需要处理长文档的系统时,必须考虑这一限制。

1.3 QKV矩阵计算详解

1.3.1 为什么需要Q、K、V三个矩阵?

Self-Attention中最关键的设计决策之一是引入三个独立的线性投影: Query(查询)Key(键)Value(值)

一个常见的误解是:为什么不直接用原始词向量计算相似度? 答案在于:Q、K、V投影使模型可以学习不同角色的表示空间

  • Query矩阵 W_Q:将输入投影为"我在寻找什么"的表示。 对于每个token,Query表示该token在寻找的信息类型;
  • Key矩阵 W_K:将输入投影为"我能提供什么"的表示。 每个token的Key描述了它所包含的信息类型;
  • Value矩阵 W_V:将输入投影为"我实际包含的内容"的表示。 Value是当注意力权重确定后实际被聚合的信息。

Query和Key在同一空间中比较(计算相似度),而Value可以在不同的空间中存在。 这种分离使得模型可以独立优化"查询-匹配"的过程和"信息提取"的过程。

1.3.2 完整的矩阵计算流程

设输入矩阵为 X ∈ ℝ^{n×d_model} ,其中 n 是序列长度, d_model 是模型隐藏维度(原始论文中为512)。

三个投影矩阵的维度:

  • W_Q ∈ ℝ^{d_model × d_k}
  • W_K ∈ ℝ^{d_model × d_k}
  • W_V ∈ ℝ^{d_model × d_v}

其中原始论文中 d_k = d_v = d_model / h = 64h=8 为head数量)。

线性投影:

ini 复制代码
  Q = X · W_Q    ∈ ℝ^{n × d_k}
  K = X · W_K    ∈ ℝ^{n × d_k}
  V = X · W_V    ∈ ℝ^{n × d_v}

注意:在Encoder-Decoder的Cross-Attention中,Q来自Decoder,K和V来自Encoder输出。 这是Self-Attention与Cross-Attention的本质区别。

1.3.3 参数量分析

一个Self-Attention层包含4个可学习矩阵:W_Q、W_K、W_V,以及输出投影矩阵 W_O。 对于 d_model=512, d_k=d_v=64, h=8 的配置:

  • 每个head的W_Q: 512×64 = 32,768个参数
  • 每个head的W_K: 512×64 = 32,768个参数
  • 每个head的W_V: 512×64 = 32,768个参数
  • 8个head共计: 3 × 8 × 32,768 = 786,432个参数
  • 输出投影W_O: 512×512 = 262,144个参数
  • 单层Self-Attention总计约100万参数

现代LLM的规模

GPT-3(175B参数)使用 d_model=12288, 96个Attention Head, 每个head的d_k=128。单个Attention层的参数量约为4×12288²≈604M。 96层Transformer共计约580亿参数仅在Attention层, 其余参数在FFN层(通常占总参数量的约2/3)。

1.4 Scaled Dot-Product Attention

1.4.1 完整公式

Scaled Dot-Product Attention的核心公式:

scss 复制代码
  Attention(Q, K, V) = softmax( Q·K^T / √d_k ) · V

分步拆解:

  1. 计算原始注意力分数scores = Q · K^T,得到形状为 n×n 的矩阵, scores[i][j] 表示位置 i 的Query与位置 j 的Key的内积;
  2. 缩放scaled_scores = scores / √d_k
  3. 可选掩码 :在Causal/Masked Attention中,将上三角部分设为 -∞
  4. Softmax归一化weights = softmax(scaled_scores),沿最后一维计算,得到行和为1的权重矩阵;
  5. 加权聚合output = weights · V, 输出形状为 n × d_v

1.4.2 缩放因子 √d_k 的必要性

这是一个容易被忽略但至关重要的设计。当 d_k 较大时(如64、128), Q·K^T 的内积值会随维度线性增长(方差为 d_k ,标准差为 √d_k)。

如果不进行缩放,高维向量的内积会产生极大的数值, 导致Softmax的输出集中在某个极值处(梯度几乎为零的"饱和区域"), 反向传播时梯度消失,训练失败。

数值验证

假设Q和K的每个分量独立地从均值为0、方差为1的分布中采样。 那么内积 q·k = Σ q_i·k_i 的期望为0,方差为 d_k 。 除以 √d_k 后,内积的方差归一化为1, Softmax的输入保持在数值稳定的范围内。

1.4.3 Causal Masking:自回归的数学保证

在语言模型(GPT系列)的训练中,模型需要预测下一个token, 因此必须确保位置 i 的输出只依赖于位置 1..i 的输入, 而不能"看到未来"。

通过在Softmax之前将注意力矩阵的上三角部分设为 -∞

ini 复制代码
# mask[i][j] = True 表示位置i不能关注位置j(j > i时)
mask = torch.triu(torch.ones(n, n), diagonal=1).bool()
scores = scores.masked_fill(mask, float('-inf'))
weights = torch.softmax(scores, dim=-1)

经过Softmax后,-∞ 的位置权重变为0, 有效地实现了因果掩码(Causal Mask)。

1.5 Multi-Head Attention机制

1.5.1 设计动机

单头注意力(Single-Head Attention)在每次前向传播中, 对于一个Query只计算一种相似度模式。然而,自然语言中的依赖关系是多维度的:

  • 语法依赖(主谓关系、动宾关系)
  • 语义关联(同义词、上位词关系)
  • 指代关系(代词与先行词)
  • 长距离依赖(从句修饰)

Multi-Head Attention (多头注意力)通过并行运行 h 个独立的注意力"头",让模型在 h 个不同的子空间中 同时学习不同类型的依赖关系,然后将结果拼接并通过线性变换整合。

1.5.2 完整计算流程

Multi-Head Attention公式:

scss 复制代码
  MultiHead(Q, K, V) = Concat(head_1, ..., head_h) · W_O

  其中 head_i = Attention(Q · W_Q^i, K · W_K^i, V · W_V^i)

每个head i 拥有独立的投影矩阵 W_Q^i ∈ ℝ^{d_model × d_k}W_K^i ∈ ℝ^{d_model × d_k}W_V^i ∈ ℝ^{d_model × d_v}

整个计算可以分为4个步骤:

  1. 并行投影 :将输入X分别投影到 h 个子空间, 得到 h 组Q、K、V;
  2. 并行注意力计算:每个head独立执行Scaled Dot-Product Attention;
  3. 拼接 :将 h 个head的输出在最后一维拼接, 得到形状为 n × (h · d_v) 的矩阵;
  4. 输出投影 :通过 W_O ∈ ℝ^{(h·d_v) × d_model} 将拼接结果映射回 d_model 维。

1.5.3 实现效率:批量矩阵乘法

在实际实现中,h 个head的Q、K、V投影不是串行计算的, 而是通过一次大型矩阵乘法并行完成:

ini 复制代码
# 等价于h个独立投影,但更高效
# W_Q的形状为 [d_model, h*d_k],一次计算所有head的Q
Q_all = X @ W_Q  # [batch, n, h*d_k]
Q_all = Q_all.view(batch, n, h, d_k).transpose(1, 2)  # [batch, h, n, d_k]

通过重塑(reshape)和转置(transpose),可以将批处理维度和head维度合并, 用单次批量矩阵乘法(torch.bmm)完成所有head的注意力计算, 充分利用GPU的并行能力。

1.5.4 可视化:各head学到了什么?

Transformer论文的后续研究(Vig 2019, Clark et al. 2019) 通过可视化注意力权重揭示了不同head的专业化分工:

  • 部分head专门追踪直接前驱词(位置偏移为1的注意力模式);
  • 部分head专注于句法关系(如动词对应宾语);
  • 部分head捕捉共指关系;
  • 部分head关注句子边界。

注意力头的可解释性

虽然注意力可视化提供了一定的可解释性线索,但研究表明(Jain & Wallace 2019) 注意力权重并不直接等价于模型的"推理路径"。 高注意力权重的位置并不一定是对最终预测最重要的。 可解释性工具(如Integrated Gradients、SHAP)提供了更可靠的特征归因方法。

1.6 位置编码(Positional Encoding)

1.6.1 位置无关性问题

Self-Attention本质上是一种**置换不变(Permutation Invariant)**操作: 将输入序列随机打乱后,每个token的注意力输出(仅依赖权重值,不依赖顺序) 在数学上是相同的。这意味着如果不显式引入位置信息, Transformer无法区分 "Dog bites man" 和 "Man bites dog"。

1.6.2 正弦-余弦位置编码(Sinusoidal PE)

原始Transformer论文提出了一种固定的、基于三角函数的位置编码:

正弦-余弦位置编码公式:

scss 复制代码
  PE(pos, 2i)   = sin(pos / 10000^(2i / d_model))
  PE(pos, 2i+1) = cos(pos / 10000^(2i / d_model))

其中 pos 是序列中的绝对位置(0到n-1), i 是编码维度的索引(0到d_model/2 - 1)。

这种编码的精妙之处在于:

  • 相对位置可计算 :对于任意固定的偏移量 kPE(pos+k) 可以表示为 PE(pos) 的线性变换, 使模型可以学习相对位置关系;
  • 不同频率编码不同尺度 :低维度(小 i)用高频正弦波编码细粒度位置, 高维度用低频正弦波编码粗粒度位置,类似傅里叶分析中的多尺度表示;
  • 外推到更长序列:理论上可以处理训练时未见过的更长序列 (尽管实践中性能会下降)。

1.6.3 可学习位置编码(Learned PE)

BERT和GPT系列采用可学习位置编码: 每个位置对应一个可训练的嵌入向量,在训练过程中通过梯度下降优化。

实现极其简单:

ini 复制代码
self.position_embedding = nn.Embedding(max_seq_len, d_model)
# 使用时:
positions = torch.arange(seq_len, device=x.device).unsqueeze(0)
x = x + self.position_embedding(positions)

可学习PE的缺点是无法外推到超过训练时最大序列长度的位置, 因为没有对应的嵌入向量。

1.6.4 旋转位置编码(RoPE)

RoPE(Rotary Position Embedding) (Su et al., 2021) 是现代LLM(LLaMA、GPT-NeoX、PaLM 2)广泛采用的位置编码方案。 其核心思想是:将位置信息编码为Query和Key向量在复数空间中的旋转, 使得内积 q_pos · k_pos' 只依赖于相对位置 (pos - pos')

RoPE的核心变换:

ini 复制代码
  q_m = R_Θ,m · q    (将Query旋转θ角,θ与位置m相关)
  k_n = R_Θ,n · k    (将Key旋转θ角,θ与位置n相关)

  q_m · k_n = (R_Θ,m·q)^T · (R_Θ,n·k)
            = q^T · R_Θ,(n-m) · k    (只依赖相对位置m-n)

为什么现代LLM偏好RoPE?

RoPE的关键优势是相对位置感知:注意力分数天然只依赖相对位置, 这对于理解自然语言("前一个词"、"句子开头"等相对位置关系)更有利。 同时,通过YaRN、LongRoPE等扩展技术,RoPE可以在推理时支持比训练时更长的上下文窗口。

1.6.5 ALiBi(Attention with Linear Biases)

ALiBi(Press et al., 2021)采用完全不同的思路: 不修改词向量本身,而是在注意力分数矩阵中直接加入与相对距离成正比的负偏置:

css 复制代码
  scores[i][j] = q_i · k_j / √d_k  -  m · |i - j|

其中 m 是每个head的斜率参数(固定或可学习)。 距离越远的token受到越大的惩罚,天然实现了局部注意力的偏好。 ALiBi在MPT、BLOOM等模型中被采用,以其优秀的长度外推能力著称。

1.7 完整Transformer架构

1.7.1 架构全貌(ASCII图)

sql 复制代码
┌─────────────────────────────────────────────────────────────┐
│                    TRANSFORMER 架构                          │
├───────────────────────────┬─────────────────────────────────┤
│        ENCODER            │           DECODER               │
│    (×N_encoder 层)        │       (×N_decoder 层)           │
│                           │                                 │
│  Input Tokens             │  Output Tokens (shifted right)  │
│       │                   │        │                        │
│  ┌────▼────┐              │  ┌─────▼─────┐                 │
│  │ Embedding│              │  │ Embedding  │                 │
│  │  + PE   │              │  │   + PE    │                 │
│  └────┬────┘              │  └─────┬─────┘                 │
│       │                   │        │                        │
│  ┌────▼──────────────┐    │  ┌─────▼───────────────────┐   │
│  │  Multi-Head       │    │  │  Masked Multi-Head       │   │
│  │  Self-Attention   │    │  │  Self-Attention          │   │
│  │  (Bidirectional)  │    │  │  (Causal / Unidirectional│   │
│  └────┬──────────────┘    │  └─────┬───────────────────┘   │
│       │                   │        │                        │
│  ┌────▼────┐              │  ┌─────▼─────┐                 │
│  │Add & Norm│              │  │ Add & Norm │                 │
│  └────┬────┘              │  └─────┬─────┘                 │
│       │                   │        │                        │
│  ┌────▼──────────────┐    │  ┌─────▼───────────────────┐   │
│  │  Feed-Forward     │    │  │  Cross-Attention         │   │
│  │  Network (FFN)    │    │  │  (Q←Decoder, K/V←Encoder│   │
│  │  2×Linear + ReLU  │    │  └─────┬───────────────────┘   │
│  └────┬──────────────┘    │        │                        │
│       │                   │  ┌─────▼─────┐                 │
│  ┌────▼────┐              │  │ Add & Norm │                 │
│  │Add & Norm│              │  └─────┬─────┘                 │
│  └────┬────┘              │        │                        │
│       │  (重复N层)         │  ┌─────▼──────────────────┐   │
│       │                   │  │  Feed-Forward Network    │   │
│  ┌────▼────┐              │  └─────┬──────────────────┘    │
│  │ Encoder │              │        │                        │
│  │ Output  │──────────────┼──►  K, V                       │
│  └─────────┘              │  ┌─────▼─────┐                 │
│                           │  │ Add & Norm │                 │
│                           │  └─────┬─────┘  (重复N层)      │
│                           │        │                        │
│                           │  ┌─────▼──────────────────┐    │
│                           │  │  Linear Projection      │    │
│                           │  │  + Softmax              │    │
│                           │  └─────┬──────────────────┘    │
│                           │        │                        │
│                           │  Output Probabilities           │
│                           │  over Vocabulary                │
└───────────────────────────┴─────────────────────────────────┘

 Encoder 结构(单层):              Decoder 结构(单层):
 ┌─────────────────────┐           ┌───────────────────────┐
 │ Multi-Head           │           │ Masked Multi-Head      │
 │ Self-Attention       │           │ Self-Attention         │
 │ ↓ Add & LayerNorm   │           │ ↓ Add & LayerNorm     │
 │ Feed-Forward Network │           │ Multi-Head             │
 │ ↓ Add & LayerNorm   │           │ Cross-Attention        │
 └─────────────────────┘           │ ↓ Add & LayerNorm     │
                                    │ Feed-Forward Network  │
                                    │ ↓ Add & LayerNorm     │
                                    └───────────────────────┘

1.7.2 Encoder vs Decoder vs Encoder-Decoder

现代大型模型按架构可分为三类:

  • 纯Encoder(如BERT):双向注意力,所有位置可互相关注, 适合理解任务(分类、NER、问答);
  • 纯Decoder(如GPT系列):单向(因果)注意力, 每个位置只能关注之前的位置,适合生成任务;
  • Encoder-Decoder(如T5、BART):编码器处理输入, 解码器通过Cross-Attention读取编码器表示并生成输出,适合翻译、摘要等seq2seq任务。

1.7.3 Decoder-Only架构的主导地位

自2020年GPT-3发布以来,几乎所有最强大的基础模型(LLaMA、Claude、Gemini、Grok) 都采用Decoder-Only架构。 原因是多方面的:

  • 统一的预训练目标(Next Token Prediction)比MLM更易扩展;
  • 无需分离的Encoder-Decoder权重,参数利用更高效;
  • 通过In-Context Learning(ICL)无需微调即可适配多种任务;
  • 生成任务覆盖范围更广(包括理解任务可转化为生成格式)。

1.8 前馈网络与残差连接

1.8.1 Position-wise Feed-Forward Network(FFN)

每个Transformer层中,Multi-Head Attention后面紧跟一个 位置独立前馈网络(Position-wise FFN)

scss 复制代码
  FFN(x) = max(0, x·W_1 + b_1) · W_2 + b_2

  其中 W_1 ∈ ℝ^{d_model × d_ff},  W_2 ∈ ℝ^{d_ff × d_model}
  原始论文中 d_ff = 4 × d_model = 2048

"Position-wise"意味着FFN对序列中每个位置独立应用相同的变换(共享权重), 不在位置之间传递信息。信息在位置间的交流完全由Attention层负责。

FFN的参数量(2 × d_model × d_ff)通常占单层参数量的约2/3, 是Transformer中参数最密集的组件。 研究表明FFN层存储了大量的事实知识(factual knowledge), 可被视为"键值记忆(key-value memories)"(Geva et al., 2021)。

1.8.2 现代FFN变体:SwiGLU和GeGLU

现代LLM普遍将ReLU激活函数替换为 SwiGLU(Swish-Gated Linear Unit):

scss 复制代码
  SwiGLU(x) = Swish(x·W_gate) ⊙ (x·W_up)
  最终输出:   SwiGLU(x) · W_down

  其中 Swish(x) = x · sigmoid(x)
  ⊙ 表示逐元素乘法

LLaMA系列使用SwiGLU时,d_ff通常设为约 8/3 × d_model(而非4×), 以保持参数量不变的前提下引入门控机制,实验表明在困惑度上显著优于ReLU/GELU。

1.8.3 残差连接与Layer Normalization

残差连接(Residual Connection)(He et al., 2016) 是深度神经网络训练的关键技术,在Transformer中每个子层后都有应用:

ini 复制代码
  x = x + SubLayer(LayerNorm(x))    # Pre-LN(现代LLM的标准)
  或
  x = LayerNorm(x + SubLayer(x))    # Post-LN(原始Transformer)

Pre-LN vs Post-LN :原始论文使用Post-LN(先计算,后归一化), 但实践发现训练不稳定,需要Learning Rate Warmup。 现代LLM(LLaMA、GPT-NeoX)普遍改用Pre-LN: 先对输入归一化,再进入子层,然后与原始输入相加。 Pre-LN训练更稳定,但理论上最后一层输出未被归一化(有些模型会在最后额外加一层归一化)。

RMSNorm(Root Mean Square Layer Normalization) 是Layer Norm的简化版本,只做缩放而不做偏移(无均值中心化), 计算更高效。LLaMA、Mistral等模型广泛使用RMSNorm替代LayerNorm:

scss 复制代码
  RMSNorm(x) = x / RMS(x) · γ

  RMS(x) = √( (1/d) · Σ x_i² )    γ 是可学习的缩放参数

1.9 Python代码实现

代码示例1:从零实现Scaled Dot-Product Attention

ini 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F
import math

def scaled_dot_product_attention(
    query: torch.Tensor,   # [batch, heads, seq_q, d_k]
    key: torch.Tensor,     # [batch, heads, seq_k, d_k]
    value: torch.Tensor,   # [batch, heads, seq_k, d_v]
    mask: torch.Tensor | None = None,  # [batch, 1, seq_q, seq_k] 或 [batch, heads, seq_q, seq_k]
    dropout_p: float = 0.0,
) -> tuple[torch.Tensor, torch.Tensor]:
    """
    Scaled Dot-Product Attention的完整实现。

    返回:
        output:  [batch, heads, seq_q, d_v]
        weights: [batch, heads, seq_q, seq_k]  (可用于可视化)
    """
    d_k = query.size(-1)

    # Step 1: 计算原始注意力分数
    # [batch, heads, seq_q, d_k] × [batch, heads, d_k, seq_k] → [batch, heads, seq_q, seq_k]
    scores = torch.matmul(query, key.transpose(-2, -1))

    # Step 2: 缩放,防止维度过大导致梯度消失
    scores = scores / math.sqrt(d_k)

    # Step 3: 应用掩码(Causal Mask 或 Padding Mask)
    if mask is not None:
        # mask 中 True 的位置会被设为 -inf,Softmax后权重为0
        scores = scores.masked_fill(mask == 0, float('-inf'))

    # Step 4: Softmax 归一化
    # 在 seq_k 维度上归一化(每个 Query 位置的权重和为1)
    weights = F.softmax(scores, dim=-1)

    # 处理全 -inf 行(padding token对应的Query行):用0替换NaN
    weights = torch.nan_to_num(weights, nan=0.0)

    # Step 5: 可选的Dropout(仅在训练时)
    if dropout_p > 0.0 and torch.is_grad_enabled():
        weights = F.dropout(weights, p=dropout_p)

    # Step 6: 加权聚合 Value
    # [batch, heads, seq_q, seq_k] × [batch, heads, seq_k, d_v] → [batch, heads, seq_q, d_v]
    output = torch.matmul(weights, value)

    return output, weights

# 验证:用随机数据测试形状
if __name__ == "__main__":
    batch_size, num_heads, seq_len, d_k, d_v = 2, 8, 16, 64, 64

    Q = torch.randn(batch_size, num_heads, seq_len, d_k)
    K = torch.randn(batch_size, num_heads, seq_len, d_k)
    V = torch.randn(batch_size, num_heads, seq_len, d_v)

    # 创建Causal Mask(下三角为1,上三角为0)
    causal_mask = torch.tril(torch.ones(seq_len, seq_len)).unsqueeze(0).unsqueeze(0)

    output, weights = scaled_dot_product_attention(Q, K, V, mask=causal_mask)

    print(f"Output shape:  {output.shape}")    # [2, 8, 16, 64]
    print(f"Weights shape: {weights.shape}")   # [2, 8, 16, 16]
    print(f"Weights row sum: {weights[0, 0, :, :].sum(dim=-1)}")  # 应约为全1

代码示例2:完整的Multi-Head Attention模块

python 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F
import math

class MultiHeadAttention(nn.Module):
    """
    完整的Multi-Head Attention实现,符合"Attention Is All You Need"论文规格。
    支持 Self-Attention(Q=K=V=x)和 Cross-Attention(Q来自Decoder,K/V来自Encoder)。
    """

    def __init__(
        self,
        d_model: int,
        num_heads: int,
        dropout: float = 0.1,
        bias: bool = True,
    ):
        super().__init__()

        assert d_model % num_heads == 0, (
            f"d_model ({d_model}) must be divisible by num_heads ({num_heads})"
        )

        self.d_model = d_model
        self.num_heads = num_heads
        self.d_k = d_model // num_heads  # 每个head的维度
        self.scale = math.sqrt(self.d_k)

        # 合并所有head的投影到单个大矩阵,效率更高
        self.W_q = nn.Linear(d_model, d_model, bias=bias)
        self.W_k = nn.Linear(d_model, d_model, bias=bias)
        self.W_v = nn.Linear(d_model, d_model, bias=bias)
        self.W_o = nn.Linear(d_model, d_model, bias=bias)

        self.dropout = nn.Dropout(dropout)

        self._init_weights()

    def _init_weights(self):
        """Xavier初始化,防止初始化时注意力分数数值异常。"""
        for module in [self.W_q, self.W_k, self.W_v, self.W_o]:
            nn.init.xavier_uniform_(module.weight)
            if module.bias is not None:
                nn.init.zeros_(module.bias)

    def _split_heads(self, x: torch.Tensor) -> torch.Tensor:
        """
        将 [batch, seq, d_model] 重塑为 [batch, num_heads, seq, d_k]。
        这是Multi-Head Attention实现的关键变换。
        """
        batch_size, seq_len, _ = x.shape
        # 先重塑为 [batch, seq, num_heads, d_k]
        x = x.view(batch_size, seq_len, self.num_heads, self.d_k)
        # 转置为 [batch, num_heads, seq, d_k],使head维度紧跟batch
        return x.transpose(1, 2)

    def _merge_heads(self, x: torch.Tensor) -> torch.Tensor:
        """
        将 [batch, num_heads, seq, d_k] 重塑回 [batch, seq, d_model]。
        逆转 _split_heads 的操作。
        """
        batch_size, _, seq_len, _ = x.shape
        # 转置回 [batch, seq, num_heads, d_k]
        x = x.transpose(1, 2).contiguous()
        # 合并最后两个维度 [batch, seq, d_model]
        return x.view(batch_size, seq_len, self.d_model)

    def forward(
        self,
        query: torch.Tensor,          # [batch, seq_q, d_model]
        key: torch.Tensor,            # [batch, seq_k, d_model]
        value: torch.Tensor,          # [batch, seq_k, d_model]
        mask: torch.Tensor | None = None,  # [batch, 1, seq_q, seq_k]
        return_attention_weights: bool = False,
    ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
        batch_size = query.size(0)

        # Step 1: 线性投影 → 拆分多头
        Q = self._split_heads(self.W_q(query))  # [batch, heads, seq_q, d_k]
        K = self._split_heads(self.W_k(key))    # [batch, heads, seq_k, d_k]
        V = self._split_heads(self.W_v(value))  # [batch, heads, seq_k, d_k]

        # Step 2: Scaled Dot-Product Attention(每个head并行计算)
        scores = torch.matmul(Q, K.transpose(-2, -1)) / self.scale  # [batch, heads, seq_q, seq_k]

        if mask is not None:
            scores = scores.masked_fill(mask == 0, float('-inf'))

        attn_weights = F.softmax(scores, dim=-1)
        attn_weights = torch.nan_to_num(attn_weights, nan=0.0)
        attn_weights_dropped = self.dropout(attn_weights)

        # Step 3: 加权聚合
        context = torch.matmul(attn_weights_dropped, V)  # [batch, heads, seq_q, d_k]

        # Step 4: 合并多头 → 输出投影
        context = self._merge_heads(context)  # [batch, seq_q, d_model]
        output = self.W_o(context)            # [batch, seq_q, d_model]

        if return_attention_weights:
            return output, attn_weights  # 返回未dropout的权重用于可视化

        return output

# 使用示例:Self-Attention
if __name__ == "__main__":
    d_model, num_heads, batch, seq = 512, 8, 4, 128

    mha = MultiHeadAttention(d_model=d_model, num_heads=num_heads)
    x = torch.randn(batch, seq, d_model)

    # Causal mask:[1, 1, seq, seq]
    causal_mask = torch.tril(torch.ones(seq, seq)).unsqueeze(0).unsqueeze(0)

    output, weights = mha(x, x, x, mask=causal_mask, return_attention_weights=True)

    print(f"Input shape:   {x.shape}")        # [4, 128, 512]
    print(f"Output shape:  {output.shape}")   # [4, 128, 512]
    print(f"Weights shape: {weights.shape}")  # [4, 8, 128, 128]

    # 验证参数量
    total_params = sum(p.numel() for p in mha.parameters())
    print(f"Total params:  {total_params:,}")  # ~1,048,576

代码示例3:完整的Transformer Encoder Block(含位置编码)

python 复制代码
import torch
import torch.nn as nn
import math

class SinusoidalPositionalEncoding(nn.Module):
    """
    原始Transformer论文的正弦-余弦位置编码。
    固定(不可学习),可外推到训练时未见过的序列长度。
    """

    def __init__(self, d_model: int, max_seq_len: int = 8192, dropout: float = 0.1):
        super().__init__()
        self.dropout = nn.Dropout(dropout)

        # 预计算位置编码矩阵,注册为buffer(不参与梯度计算)
        pe = torch.zeros(max_seq_len, d_model)

        position = torch.arange(0, max_seq_len, dtype=torch.float).unsqueeze(1)  # [max_seq, 1]
        div_term = torch.exp(
            torch.arange(0, d_model, 2, dtype=torch.float) * (-math.log(10000.0) / d_model)
        )  # [d_model/2]

        pe[:, 0::2] = torch.sin(position * div_term)  # 偶数维度用sin
        pe[:, 1::2] = torch.cos(position * div_term)  # 奇数维度用cos

        pe = pe.unsqueeze(0)  # [1, max_seq, d_model]
        self.register_buffer('pe', pe)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        # x: [batch, seq, d_model]
        x = x + self.pe[:, :x.size(1), :]
        return self.dropout(x)

class FeedForwardNetwork(nn.Module):
    """
    Position-wise FFN:两层线性变换 + GELU激活。
    现代实现通常用GELU替代原始的ReLU。
    """

    def __init__(self, d_model: int, d_ff: int, dropout: float = 0.1):
        super().__init__()
        self.linear1 = nn.Linear(d_model, d_ff)
        self.linear2 = nn.Linear(d_ff, d_model)
        self.dropout = nn.Dropout(dropout)
        self.activation = nn.GELU()

        # 使用专门针对Transformer的初始化
        nn.init.xavier_uniform_(self.linear1.weight)
        nn.init.xavier_uniform_(self.linear2.weight)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return self.linear2(self.dropout(self.activation(self.linear1(x))))

class TransformerEncoderBlock(nn.Module):
    """
    标准Transformer Encoder层,使用Pre-LN(现代标准)。

    结构:
        x = x + Attention(LayerNorm(x))
        x = x + FFN(LayerNorm(x))
    """

    def __init__(
        self,
        d_model: int = 512,
        num_heads: int = 8,
        d_ff: int = 2048,
        dropout: float = 0.1,
    ):
        super().__init__()

        self.self_attention = MultiHeadAttention(
            d_model=d_model,
            num_heads=num_heads,
            dropout=dropout,
        )
        self.ffn = FeedForwardNetwork(d_model=d_model, d_ff=d_ff, dropout=dropout)

        # Pre-LN 使用 LayerNorm 在子层之前(而非之后)
        self.norm1 = nn.LayerNorm(d_model)
        self.norm2 = nn.LayerNorm(d_model)

        self.dropout = nn.Dropout(dropout)

    def forward(
        self,
        x: torch.Tensor,
        padding_mask: torch.Tensor | None = None,
    ) -> torch.Tensor:
        # Self-Attention子层(Pre-LN)
        residual = x
        x_norm = self.norm1(x)
        attn_output = self.self_attention(
            query=x_norm,
            key=x_norm,
            value=x_norm,
            mask=padding_mask,
        )
        x = residual + self.dropout(attn_output)

        # FFN子层(Pre-LN)
        residual = x
        x = residual + self.dropout(self.ffn(self.norm2(x)))

        return x

class TransformerEncoder(nn.Module):
    """
    完整的Transformer Encoder:Embedding + PE + N个Encoder Block。
    """

    def __init__(
        self,
        vocab_size: int,
        d_model: int = 512,
        num_heads: int = 8,
        num_layers: int = 6,
        d_ff: int = 2048,
        max_seq_len: int = 512,
        dropout: float = 0.1,
    ):
        super().__init__()

        self.d_model = d_model
        self.embedding = nn.Embedding(vocab_size, d_model)
        self.pos_encoding = SinusoidalPositionalEncoding(d_model, max_seq_len, dropout)

        self.layers = nn.ModuleList([
            TransformerEncoderBlock(d_model, num_heads, d_ff, dropout)
            for _ in range(num_layers)
        ])

        # 最后一层额外的LayerNorm(Pre-LN架构需要)
        self.final_norm = nn.LayerNorm(d_model)

        self._init_embeddings()

    def _init_embeddings(self):
        nn.init.normal_(self.embedding.weight, mean=0.0, std=self.d_model ** -0.5)

    def forward(
        self,
        input_ids: torch.Tensor,           # [batch, seq]
        attention_mask: torch.Tensor | None = None,  # [batch, seq],1=有效,0=padding
    ) -> torch.Tensor:
        # 将 padding mask 扩展到 [batch, 1, 1, seq] 供 Attention 使用
        if attention_mask is not None:
            extended_mask = attention_mask.unsqueeze(1).unsqueeze(2)
        else:
            extended_mask = None

        # Embedding + Positional Encoding
        x = self.embedding(input_ids) * math.sqrt(self.d_model)
        x = self.pos_encoding(x)

        # 逐层处理
        for layer in self.layers:
            x = layer(x, padding_mask=extended_mask)

        return self.final_norm(x)

# 完整使用示例
if __name__ == "__main__":
    # BERT-Base规格的Encoder
    encoder = TransformerEncoder(
        vocab_size=30522,
        d_model=768,
        num_heads=12,
        num_layers=12,
        d_ff=3072,
        max_seq_len=512,
        dropout=0.1,
    )

    # 统计参数量
    total_params = sum(p.numel() for p in encoder.parameters())
    print(f"Total parameters: {total_params:,}")  # ~86M(BERT-Base规格)

    # 前向传播测试
    batch_size, seq_len = 4, 128
    input_ids = torch.randint(0, 30522, (batch_size, seq_len))
    attention_mask = torch.ones(batch_size, seq_len)
    attention_mask[0, 100:] = 0  # 模拟第一个样本有padding

    output = encoder(input_ids, attention_mask)
    print(f"Input shape:  {input_ids.shape}")   # [4, 128]
    print(f"Output shape: {output.shape}")       # [4, 128, 768]
    print(f"Output norm (mean): {output.norm(dim=-1).mean():.4f}")

性能优化提示

上面的实现以清晰性为优先,不是最优的生产实现。 在实际工程中,应使用 torch.nn.functional.scaled_dot_product_attention(PyTorch 2.0+), 它会自动选择FlashAttention等融合内核,内存效率提升2-4倍,速度提升2-8倍。 对于推理,可进一步使用KV Cache(在自回归生成中复用历史K/V,避免重复计算)。

1.10 架构对比

对比表1:序列模型架构横向对比

特性 RNN / LSTM CNN(1D卷积) Transformer 现代LLM变体(如LLaMA)
并行训练 不支持(顺序依赖) 支持(局部并行) 完全并行 完全并行
长距离依赖 受梯度消失限制(LSTM有限改善) 取决于感受野大小 直接(O(1)步) 直接(扩展上下文窗口)
计算复杂度 O(n·d²) O(n·k·d²)(k=卷积核大小) O(n²·d)(注意力二次方) O(n²·d)(可用Flash优化)
位置感知 内建(隐式顺序) 内建(卷积位移不变性) 需显式PE(可学习/固定) RoPE / ALiBi(相对位置)
内存占用 O(n)(隐藏状态) O(n·d) O(n²)(注意力矩阵) O(n²) → O(n)(FlashAttention)
推理(逐步生成) O(d²)每步(状态复用) 需要固定窗口 O(n·d²)每步(KV Cache后O(d²)) O(d²)每步(KV Cache)
可扩展性 难以扩展到超大模型 有限 出色(证明于GPT-3等) 极佳(数千亿参数)
代表模型 ELMo, seq2seq TextCNN, ByteNet BERT, T5, GPT-2 LLaMA, Claude, GPT-4

对比表2:位置编码方案对比

方案 类型 可学习 外推能力 相对位置感知 实现复杂度 采用模型
Sinusoidal PE 绝对位置 理论可外推(实践差) 间接(线性变换) 极简 原始Transformer
Learned PE 绝对位置 不可外推 简单 BERT, GPT-2
Relative PE(Shaw) 相对位置 有限外推 直接 中等 Transformer-XL
RoPE 相对位置(旋转) 否(固定旋转) 可扩展(YaRN/LongRoPE) 天然(内积即相对) 中等 LLaMA, GPT-NeoX, PaLM
ALiBi 相对位置(偏置) 否(固定斜率) 出色 直接(距离惩罚) 简单 MPT, BLOOM
NoPE(无PE) --- --- 依赖因果掩码 无显式位置 最简 部分研究模型

图例:绿色表示优势,橙色表示中等,红色表示劣势。

1.11 关键要点总结

Self-Attention核心
  • 本质是加权聚合:Attention(Q,K,V) = softmax(QK^T/√d_k)V
  • O(n²·d)复杂度是主要瓶颈,限制了超长序列处理
  • Q/K/V三路投影让模型分离"查询"与"内容"表示空间
  • 缩放因子√d_k防止大维度下Softmax饱和
Multi-Head Attention
  • h个head并行捕捉不同语义维度的依赖关系
  • 通过reshape+transpose实现高效并行计算
  • 输出投影W_O整合多头信息
  • 实践中不同head专业化于不同语言现象
位置编码选择
  • RoPE已成为现代LLM的事实标准(LLaMA, GPT-NeoX)
  • ALiBi在长度外推上表现优异
  • 可学习PE简单但无法处理超出训练长度的序列
  • 正弦PE历史重要但实践中已少用
架构组件
  • Pre-LN比Post-LN训练更稳定,已成现代标准
  • RMSNorm替代LayerNorm,更高效(LLaMA等)
  • SwiGLU/GeGLU替代ReLU,FFN性能显著提升
  • 残差连接是深度Transformer训练的关键保障
工程实践要点
  • 使用torch.nn.functional.scaled_dot_product_attention(PyTorch 2.0+)自动选择FlashAttention
  • 推理时启用KV Cache,将每步复杂度从O(n·d²)降至O(d²)
  • 批量矩阵乘法(BMM)是Multi-Head实现的效率关键
  • Causal Mask通过masked_fill(-inf)实现自回归约束
架构演进方向
  • FlashAttention:IO感知实现,HBM峰值显存O(n²)→O(n)(FLOP仍O(n²))
  • GQA/MQA:分组查询注意力,减少KV Cache内存占用
  • Sparse Attention(Longformer, BigBird):O(n²)→O(n)
  • 状态空间模型(Mamba):对某些任务可替代Attention

常见误区

误区1 :注意力权重直接代表重要性。实际上,高注意力权重未必对应高特征归因(Jain & Wallace, 2019)。

误区2 :Transformer一定优于RNN。对于流式处理、超长序列(>100K tokens)或内存受限场景, RNN/SSM(状态空间模型)仍有优势。

误区3:层数越多越好。过多层数在没有足够数据和算力的情况下会导致过拟合, 需要根据任务规模平衡深度与宽度。

延伸阅读

  • Vaswani et al. (2017) --- "Attention Is All You Need" --- 原始Transformer论文
  • Devlin et al. (2018) --- "BERT: Pre-training of Deep Bidirectional Transformers"
  • Su et al. (2021) --- "RoFormer: Enhanced Transformer with Rotary Position Embedding"
  • Dao et al. (2022) --- "FlashAttention: Fast and Memory-Efficient Exact Attention"
  • Geva et al. (2021) --- "Transformer Feed-Forward Layers Are Key-Value Memories"
  • Press et al. (2021) --- "Train Short, Test Long: Attention with Linear Biases (ALiBi)"
相关推荐
生信大杂烩1 小时前
Xenium H&E空间原位可视化——细胞轮廓、基因表达与转录本可视化
python·算法·数据分析
linx2952 小时前
第七章 · 标准库容器、算法与 ranges
c语言·开发语言·数据结构·c++·算法
钓鱼的肝2 小时前
csp-j-s总结(4)
c++·经验分享·笔记·算法
橘子汽水1682 小时前
Leetcode 763,45 划分字母区间 跳跃游戏II
数据结构·算法·leetcode
天天喝旺仔2 小时前
Git 内部原理深度解析:从 blob/tree/commit 对象到 packfile 与垃圾回收
数据结构·数据库·git·算法·哈希
AIGCmagic社区2 小时前
KITTI AbsRel从6.5压到5.4,Marigold V2用一张32GB卡把编辑DiT收成单步深度估计
人工智能·算法·aigc·ai多模态
青少儿编程课堂2 小时前
多源最短路与最小环(Floyd 算法图论解析)
c++·python·算法·bfs·信息学竞赛
6Hzlia3 小时前
【Classic 150 刷题计划】 LeetCode 228. 汇总区间 | C++ 锚点游标与断点检测法
c++·算法·leetcode
Edward The Bunny3 小时前
Leetcode Hot 100
数据结构·算法