Skip to content

Transformer 架构 ​

标签
AI/llm
字数
3323 字
阅读时间
14 分钟

Transformer 由谷歌团队 2017 年提出(Vaswani et al., Attention is all you need, NeurIPS 2017)。它完全抛弃循环结构,只靠**注意力(Attention)**捕捉序列内依赖,从而做到真正的并行计算——RNN 第 t 步必须等第 t−1 步,无法并行,这是它被替换掉的直接原因。

Encoder-Decoder 整体结构 ​

最初是为端到端机器翻译设计的,宏观上分两半:

部分职责
编码器(Encoder)「理解」输入的整个句子,为每个词元生成富含上下文信息的向量表示
解码器(Decoder)「生成」目标句子,参考自己已生成的前文,并「咨询」编码器的理解结果来生成下一个词

编码器层 EncoderLayer 的结构:多头自注意力 → Add & Norm → 前馈网络 → Add & Norm。 解码器层 DecoderLayer 多一个子层:掩码多头自注意力 → Add & Norm → 交叉注意力(Q 来自解码器自己,K/V 来自编码器输出)→ Add & Norm → 前馈网络 → Add & Norm。

python
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)代表词元本身携带的「内容」或「信息」

三者由原始词嵌入乘以三个可学习权重矩阵 WQ,WK,WV 得到。计算过程:

  1. 为每个词生成 Q,K,V
  2. 相关性得分:用当前词的 Q 与所有词(含自己)的 K 做点积
  3. 稳定化与归一化:除以缩放因子 dk(dk 是 K 的维度)防止梯度过小,再 Softmax 成总和为 1 的权重
  4. 加权求和:权重分别乘以各词的 V 再相加,得到融合全局上下文的新表示

公式:

Attention(Q,K,V)=softmax(QKTdk)V

自注意力的计算是四步,前两步决定「看谁」,后两步决定「拿到什么」。

自注意力的四步(以「读到 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 在维度上切成 h 份,每份独立做一次单头注意力,最后拼接再过一个线性变换整合。相当于让 h 个「专家」从不同表示子空间审视同一个句子。

python
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

多头注意力是把 dmodel 切成 h 份分别做注意力,再拼回来。

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 就是在挡这个)。

为什么是它胜出:三个维度上的对照 ​

「并行」只是表面说法。三个量把自注意力和循环/卷积放在一起比,才是它胜出的完整理由。设 n 为序列长度、d 为表示维度、k 为卷积核宽度:

层类型每层计算复杂度顺序操作数最大路径长度
自注意力O(n2⋅d)O(1)O(1)
循环层(RNN/LSTM)O(n⋅d2)O(n)O(n)
卷积层O(k⋅n⋅d2)O(1)O(logk⁡n)

三列各回答一个不同的问题,要分开读:

一、计算复杂度:注意力不是无条件更便宜。

O(n2⋅d)vsO(n⋅d2)

只有当 n<d 时,自注意力才比循环层更便宜(可以先把 n2d 与 nd2 约掉一个 nd,剩下 n 与 d 的对比)。

这条判据很重要:它说明注意力的二次项在序列长度超过表示维度时才成为瓶颈。2017 年的机器翻译任务里 n 通常在几十到几百、d 是 512 —— 正好落在注意力更便宜的那一侧。

这也预告了后来的事:当上下文从几十上百扩到 32K、128K 时,n 远远超过 d,二次项就成了主要矛盾——这才是 FlashAttention、稀疏注意力、滑窗注意力这一系列工作的动机来源。

二、顺序操作数:这是「能不能并行」的精确定义。

自注意力是 O(1)——所有位置的计算互不依赖,一次矩阵乘全部算完。循环层是 O(n)——第 t 步必须等第 t−1 步。这一列的差距是训练能不能吃满算力的问题(见 03-Decoder-Only 与自回归 里「训练并行、生成串行」那条不对称是怎么来的)。

三、最大路径长度:这是「长程依赖好不好学」的精确定义。

自注意力里任意两个位置直接相连,路径长度恒为 1;循环层要走 O(n) 步;卷积层要 O(logk⁡n) 层堆叠。

路径越长,梯度回传经过的连乘越多,越容易衰减 —— 这正是 RNN 的梯度消失的根源。路径长度从 O(n) 压到 O(1),是 Transformer 能解决长程依赖的根本原因,而不是「注意力比循环更聪明」。

三列合起来:算力上打平(在 n<d 时还占优)、并行上碾压、长程依赖上从根上解决。这才是一个架构替换另一个架构的完整账。

三列回答三个不同的问题,混着读会得出错误的结论。

                  自注意力      循环层        卷积层
每层计算复杂度     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 次,但所有位置共享同一组权重,既保留独立加工能力,又大幅减少参数量。

FFN(x)=max(0,xW1+b1)W2+b2

通常 d_ff = 4 * d_model:先把维度放大,过 ReLU,再映射回 d_model。这种「先扩大再缩小」被认为有助于学到更丰富的特征表示。

残差连接与层归一化 ​

每个子模块都被 Add & Norm 包裹,作用有两个:

操作解决的问题机制
残差连接(Add)梯度消失Output=x+Sublayer(x),反向传播时梯度可绕过子模块直接前传
层归一化(Norm)内部协变量偏移(Internal Covariate Shift)对单个样本的所有特征归一化到均值 0、方差 1,使每层输入分布稳定

位置编码 ​

自注意力本身不含任何位置信息——对它来说 agent learns 和 learns agent 完全等价。位置编码(Positional Encoding)解决这个:给每个词元的嵌入向量额外加一个代表绝对/相对位置的「位置向量」。

它的关键特点是不通过学习得到,而是用固定数学公式直接算:

PE(pos,2i)=sin⁡(pos100002i/dmodel),PE(pos,2i+1)=cos⁡(pos100002i/dmodel)

其中 pos 是词元在序列中的位置,i 是位置向量的维度索引(0 到 dmodel/2),dmodel 是词嵌入维度。偶数维用 sin、奇数维用 cos。

这样即使两个词元同叫 agent、嵌入完全相同,由于位置不同,加上不同的位置编码后,输入到模型的向量就变得独一无二。

注意这是 2017 年原版的做法(绝对位置编码,加到嵌入上)。这一层后来演进了很长一段:绝对 → 相对 → RoPE → 长上下文外推(位置插值 / NTK-aware / YaRN)→ ALiBi。今天主流模型用的 RoPE 与这里的做法差别很大——它不对嵌入做加法,而是旋转注意力里的 Q、K。完整演进见 09-位置编码。

python
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)

相关 ​

参考 ​

  • 《Hello-Agents》第三章 §3.1.2
  • Vaswani, A., et al. Attention is all you need. NeurIPS, 2017.

贡献者 ​

文件历史 ​