Transformer 架构
Transformer 由谷歌团队 2017 年提出(Vaswani et al., Attention is all you need, NeurIPS 2017)。它完全抛弃循环结构,只靠**注意力(Attention)**捕捉序列内依赖,从而做到真正的并行计算——RNN 第
Encoder-Decoder 整体结构
最初是为端到端机器翻译设计的,宏观上分两半:
| 部分 | 职责 |
|---|---|
| 编码器(Encoder) | 「理解」输入的整个句子,为每个词元生成富含上下文信息的向量表示 |
| 解码器(Decoder) | 「生成」目标句子,参考自己已生成的前文,并「咨询」编码器的理解结果来生成下一个词 |
编码器层 EncoderLayer 的结构:多头自注意力 → Add & Norm → 前馈网络 → Add & Norm。 解码器层 DecoderLayer 多一个子层:掩码多头自注意力 → Add & Norm → 交叉注意力(Q 来自解码器自己,K/V 来自编码器输出)→ Add & Norm → 前馈网络 → Add & Norm。
class EncoderLayer(nn.Module):
def forward(self, x, mask):
attn_output = self.self_attn(x, x, x, mask)
x = self.norm1(x + self.dropout(attn_output))
ff_output = self.feed_forward(x)
x = self.norm2(x + self.dropout(ff_output))
return x原始 Transformer 是完整的编码器—解码器,两半的层结构不同,差别集中在解码器多出来的那个子层。
原始 Transformer:编码器 N 层 + 解码器 N 层
输入词元 + 位置编码
│
▼
┌─ 编码器层 × N ────────────────────────────┐
│ 多头自注意力(双向:每个位置看到整句) │
│ │ │
│ Add & Norm │
│ │ │
│ 前馈网络 FFN → Add & Norm │
└──────┬───────────────────────────────────┘
│ 编码器输出(后面当 K、V 用)
▼
┌─ 解码器层 × N ────────────────────────────┐
│ 掩码多头自注意力(因果:只看左侧已生成的) │
│ │ │
│ Add & Norm │
│ │ │
│ 交叉注意力 Q ← 解码器自己 │
│ K、V ← 编码器输出 │
│ │ │
│ Add & Norm → 前馈网络 FFN → Add & Norm │
└──────┬───────────────────────────────────┘
▼
线性层 + Softmax ──▶ 目标词表上的分布编码器与解码器的差别只有一处加一个掩码:双向 vs 因果,以及解码器多一次「咨询编码器」。
自注意力:Q、K、V
以句子 The agent learns because **it** is intelligent. 为例:读到 it 时,要理解它的指代,就得把注意力放到前面的 agent 上。自注意力就是对这个过程的数学建模——处理每个词时兼顾所有其他词,并分配不同的注意力权重。
每个输入词元向量被投影成三个可学习的角色:
| 角色 | 含义 |
|---|---|
| 查询(Query, Q) | 代表当前词元,它正在主动「查询」其他词元以获取信息 |
| 键(Key, K) | 代表句子中可被查询的词元的「标签」或「索引」 |
| 值(Value, V) | 代表词元本身携带的「内容」或「信息」 |
三者由原始词嵌入乘以三个可学习权重矩阵
- 为每个词生成
- 相关性得分:用当前词的
与所有词(含自己)的 做点积 - 稳定化与归一化:除以缩放因子
( 是 的维度)防止梯度过小,再 Softmax 成总和为 1 的权重 - 加权求和:权重分别乘以各词的
再相加,得到融合全局上下文的新表示
公式:
自注意力的计算是四步,前两步决定「看谁」,后两步决定「拿到什么」。
自注意力的四步(以「读到 it,要回看 agent」为例)
① 投影角色 x ──×W^Q──▶ Q 当前这个位置在查询什么
x ──×W^K──▶ K 每个位置可被查询的「标签」
x ──×W^V──▶ V 每个位置携带的「内容」
② 相关性得分 Q · Kᵀ 每个位置 × 每个位置 → n×n 分数矩阵
(含自己,所以对角线也在里面)
③ 稳定化+归一 ÷ √d_k → softmax(dim=-1)
缩放的作用:d_k 大时点积的量级跟着变大,不缩放会让 softmax
落进梯度极小的饱和区
归一化后每一行的权重和为 1
④ 加权求和 attn_probs × V 每个位置得到融合全局上下文的新表示
Attention(Q, K, V) = softmax(QKᵀ / √d_k) V从单头到多头
只做一次注意力,模型可能只学会关注一种关系(处理 it 时只学会关注主语)。语言里的关系是多种并存的——指代、时态、从属。多头注意力的做法是:把 Q、K、V 在维度上切成
class MultiHeadAttention(nn.Module):
def __init__(self, d_model, num_heads):
assert d_model % num_heads == 0
self.d_model, self.num_heads = d_model, 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 scaled_dot_product_attention(self, Q, K, V, mask=None):
attn_scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
if mask is not None:
attn_scores = attn_scores.masked_fill(mask == 0, -1e9)
attn_probs = torch.softmax(attn_scores, dim=-1)
return torch.matmul(attn_probs, V)
def split_heads(self, x):
batch_size, seq_length, d_model = x.size()
return x.view(batch_size, seq_length, self.num_heads, self.d_k).transpose(1, 2)
def combine_heads(self, x):
batch_size, num_heads, seq_length, d_k = x.size()
return x.transpose(1, 2).contiguous().view(batch_size, seq_length, self.d_model)
def forward(self, Q, K, V, mask=None):
Q = self.split_heads(self.W_q(Q))
K = self.split_heads(self.W_k(K))
V = self.split_heads(self.W_v(V))
attn_output = self.scaled_dot_product_attention(Q, K, V, mask)
return self.W_o(self.combine_heads(attn_output))两个实现细节值得记:
d_model必须能被num_heads整除,否则split_heads的 reshape 不成立- 掩码用
masked_fill(mask == 0, -1e9):把要屏蔽的位置填成极小负数,过 Softmax 后概率趋近 0
多头注意力是把
d_model = 512、h = 8 的例子
x(512 维)
│ W^Q / W^K / W^V 三个线性变换
▼
┌─────────┬─────────┬─────┬─────────┐
│ 头 1 │ 头 2 │ ... │ 头 h │ 每个头 d_k = d_model / h = 64
│ 64 维 │ 64 维 │ │ 64 维 │
└────┬────┴────┬────┴─────┴────┬────┘
│ │ │
▼ ▼ ▼
各自独立做一次 scaled dot-product attention(头与头之间不通信)
│ │ │
└─────────┴───────┬───────┘
▼
拼接回 512 维
│
▼
线性层 W^O ──▶ 输出d_model 必须能被 num_heads 整除,否则 split_heads 里的 reshape 不成立(代码里那条 assert d_model % num_heads == 0 就是在挡这个)。
为什么是它胜出:三个维度上的对照
「并行」只是表面说法。三个量把自注意力和循环/卷积放在一起比,才是它胜出的完整理由。设
| 层类型 | 每层计算复杂度 | 顺序操作数 | 最大路径长度 |
|---|---|---|---|
| 自注意力 | |||
| 循环层(RNN/LSTM) | |||
| 卷积层 |
三列各回答一个不同的问题,要分开读:
一、计算复杂度:注意力不是无条件更便宜。
只有当
这条判据很重要:它说明注意力的二次项在序列长度超过表示维度时才成为瓶颈。2017 年的机器翻译任务里
通常在几十到几百、 是 512 —— 正好落在注意力更便宜的那一侧。
这也预告了后来的事:当上下文从几十上百扩到 32K、128K 时,
二、顺序操作数:这是「能不能并行」的精确定义。
自注意力是
三、最大路径长度:这是「长程依赖好不好学」的精确定义。
自注意力里任意两个位置直接相连,路径长度恒为 1;循环层要走
路径越长,梯度回传经过的连乘越多,越容易衰减 —— 这正是 RNN 的梯度消失的根源。路径长度从
压到 ,是 Transformer 能解决长程依赖的根本原因,而不是「注意力比循环更聪明」。
三列合起来:算力上打平(在
三列回答三个不同的问题,混着读会得出错误的结论。
自注意力 循环层 卷积层
每层计算复杂度 O(n²·d) O(n·d²) O(k·n·d²)
顺序操作数 O(1) O(n) O(1)
最大路径长度 O(1) O(n) O(log_k n)
① 计算复杂度:交叉点在 n = d
n < d ──▶ 注意力更便宜。2017 年的机器翻译任务 n 在几十到几百、d = 512,
正落在这一侧
n > d ──▶ 二次项成为主要矛盾。32K / 128K 上下文就在这一侧,FlashAttention、
稀疏注意力、滑窗注意力都在解这一项
② 顺序操作数:这一列是「能不能并行」的精确定义
O(1) ──▶ 一次矩阵乘算完所有位置 ──▶ 训练能吃满算力
O(n) ──▶ 第 t 步必须等第 t−1 步
③ 最大路径长度:这一列是「长程依赖好不好学」的精确定义
O(1) ──▶ 任意两个位置直接相连,梯度回传只经过一次
O(n) ──▶ 梯度连乘 n 次 ──▶ 衰减(RNN 梯度消失的根源)
O(log_k n) ──▶ 卷积层堆叠出来的距离逐位置前馈网络 FFN
每个 Encoder / Decoder 层里,多头注意力之后都跟一个逐位置前馈网络(Position-wise Feed-Forward Network)。分工是:注意力层从整个序列「动态聚合」信息,前馈网络从聚合后的信息里提取更高阶特征。
「逐位置」指它独立作用于每一个词元向量——长度为 seq_len 的序列实际会调用 seq_len 次,但所有位置共享同一组权重,既保留独立加工能力,又大幅减少参数量。
通常 d_ff = 4 * d_model:先把维度放大,过 ReLU,再映射回 d_model。这种「先扩大再缩小」被认为有助于学到更丰富的特征表示。
残差连接与层归一化
每个子模块都被 Add & Norm 包裹,作用有两个:
| 操作 | 解决的问题 | 机制 |
|---|---|---|
| 残差连接(Add) | 梯度消失 | |
| 层归一化(Norm) | 内部协变量偏移(Internal Covariate Shift) | 对单个样本的所有特征归一化到均值 0、方差 1,使每层输入分布稳定 |
位置编码
自注意力本身不含任何位置信息——对它来说 agent learns 和 learns agent 完全等价。位置编码(Positional Encoding)解决这个:给每个词元的嵌入向量额外加一个代表绝对/相对位置的「位置向量」。
它的关键特点是不通过学习得到,而是用固定数学公式直接算:
其中
这样即使两个词元同叫 agent、嵌入完全相同,由于位置不同,加上不同的位置编码后,输入到模型的向量就变得独一无二。
注意这是 2017 年原版的做法(绝对位置编码,加到嵌入上)。这一层后来演进了很长一段:绝对 → 相对 → RoPE → 长上下文外推(位置插值 / NTK-aware / YaRN)→ ALiBi。今天主流模型用的 RoPE 与这里的做法差别很大——它不对嵌入做加法,而是旋转注意力里的 Q、K。完整演进见 09-位置编码。
class PositionalEncoding(nn.Module):
def __init__(self, d_model: int, dropout: float = 0.1, max_len: int = 5000):
super().__init__()
self.dropout = nn.Dropout(p=dropout)
position = torch.arange(max_len).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
pe = torch.zeros(max_len, d_model)
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
self.register_buffer('pe', pe.unsqueeze(0)) # buffer 不是参数,但会随模型 to(device)
def forward(self, x):
x = x + self.pe[:, :x.size(1)]
return self.dropout(x)相关
- 01-语言模型演进:N-gram 到 RNN —— Transformer 替换掉的那套循环结构
- 03-Decoder-Only 与自回归 —— 从完整架构砍到只剩解码器
参考
- 《Hello-Agents》第三章 §3.1.2
- Vaswani, A., et al. Attention is all you need. NeurIPS, 2017.
YJ