Transformer架构的数学原理

Transformer 架构的数学原理

神经网络

2017 年,Google 在论文《Attention Is All You Need》中提出了 Transformer 架构,彻底改变了自然语言处理领域。本文将从数学角度深入解析 Transformer 的核心——自注意力机制(Self-Attention)。

一、从 RNN 到 Transformer

1.1 RNN 的局限性

传统的 RNN(循环神经网络)存在以下问题:

  • 梯度消失/爆炸:长序列训练困难
  • 顺序计算:无法并行化,训练速度慢
  • 长距离依赖:对相隔较远的位置信息捕捉困难

1.2 Transformer 的革命性突破

Transformer 通过注意力机制彻底解决了上述问题:

  1. 并行计算:摆脱序列依赖,可高效并行训练
  2. 直接建模:任意位置间可直接建立联系
  3. 可解释性:注意力权重可视化

二、自注意力机制(Self-Attention)

2.1 核心思想

Self-Attention 的核心是:通过 Query、Key、Value 三个向量,计算序列内部每个位置对其他位置的注意力权重。

2.2 数学推导

设输入序列为 $X = (x_1, x_2, …, x_n)$,每个 $x_i$ 是 d 维向量。

第一步:线性变换

$$
Q = X W_Q, \quad K = X W_K, \quad V = X W_V
$$

其中 $W_Q, W_K, W_V$ 是可学习的参数矩阵。

第二步:计算注意力分数

$$
\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right) V
$$

其中 $\sqrt{d_k}$ 是缩放因子,用于防止点积过大导致梯度消失。

2.3 多头注意力(Multi-Head Attention)

为了捕捉不同类型的依赖关系,Transformer 使用多头注意力:

$$
\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, …, \text{head}_h) W^O
$$

其中每个 $\text{head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)$

为什么多头有效?

  • 不同的头可以关注不同的语义关系
  • 有的头关注句法,有的头关注语义
  • 有的头关注局部,有的头关注全局

三、位置编码(Positional Encoding)

3.1 为什么需要位置编码?

由于 Self-Attention 本身是位置无关的(对输入序列的每个位置一视同仁),需要显式加入位置信息。

3.2 正弦/余弦位置编码

Transformer 使用正弦和余弦函数生成位置编码:

$$
PE_{(pos, 2i)} = \sin\left(\frac{pos}{10000^{2i/d}}\right)
$$
$$
PE_{(pos, 2i+1)} = \cos\left(\frac{pos}{10000^{2i/d}}\right)
$$

这种编码方式的优势:

  • 任意两个位置的关系可以通过线性变换获得
  • 可以泛化到训练时未见过的序列长度

四、前馈神经网络(FFN)

每个 Transformer 块还包含两层全连接网络:

$$
\text{FFN}(x) = \max(0, xW_1 + b_1)W_2 + b_2
$$

通常是内维 expansion(一般 4 倍),然后压缩回来。

五、残差连接与层归一化

5.1 残差连接

$$
x’ = x + \text{SubLayer}(x)
$$

缓解深层网络的梯度消失问题。

5.2 层归一化

$$
y = \frac{x - \mu}{\sqrt{\sigma^2 + \epsilon}} \cdot \gamma + \beta
$$

稳定训练过程。

六、完整架构图

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
Input

Embedding + Positional Encoding

[Block 1]
├── Multi-Head Self-Attention
├── Add & Layer Norm
├── Feed Forward Network
└── Add & Layer Norm

[Block 2]
└── ... (重复 N)

LinearSoftmax

Output

七、代码实现要点

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
import torch
import torch.nn as nn
import math

class SelfAttention(nn.Module):
def __init__(self, d_model, num_heads):
super().__init__()
self.num_heads = num_heads
self.d_k = d_model // num_heads

self.W_q = nn.Linear(d_model, d_model)
self.W_k = nn.Linear(d_model, d_model)
self.W_v = nn.Linear(d_model, d_model)
self.W_o = nn.Linear(d_model, d_model)

def forward(self, x):
batch_size = x.size(0)

# 线性变换并分头
Q = self.W_q(x).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)
K = self.W_k(x).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)
V = self.W_v(x).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)

# 注意力计算
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
attn_weights = torch.softmax(scores, dim=-1)
context = torch.matmul(attn_weights, V)

# 合并多头并输出
context = context.transpose(1, 2).contiguous().view(batch_size, -1, self.num_heads * self.d_k)
return self.W_o(context)

八、总结

Transformer 的核心是自注意力机制,通过 Query-Key-Value 的方式计算序列内部的依赖关系。相比 RNN:

  • ✅ 并行计算效率高
  • ✅ 可捕捉任意距离的依赖
  • ✅ 可解释性强

这解释了为什么 Transformer 能在 NLP 领域取得巨大成功,并正在向 CV、Audio 等领域扩展。


下期预告:《注意力机制变体:从 Self-Attention 到 FlashAttention》


Transformer架构的数学原理
https://www.eternalquest.top/2026/05/28/transformer-math-principles/
作者
未竟之路上的行者
发布于
2026年5月28日
许可协议