架构演进 · 起源

Transformer(2017)

Vaswani 等,“Attention Is All You Need” · GPT-2 改了什么

对比
基座 2017 vs GPT-2 small
层
6 编码 + 6 解码 GPT-2 12
dmodel
512 GPT-2 768
头
8 GPT-2 12
dff
2,048 GPT-2 3,072
词表
约 37,000,两种语言共享 GPT-2 50,257
参数
65M GPT-2 124M

全部步骤

  1. 之前:一次一步

    2017 年之前,翻译模型是循环神经网络(RNN),一次读一个词元:每一步都依赖前一步,所以 GPU 无法并行地跑它们。自注意力用一次矩阵乘法就把每一对位置连接起来。

    RNN:n 步 · 注意力:1 步

  2. GPT-2 vs 2017 年的 Transformer

    GPT-2 只有一摞。2017 年的 Transformer 有两摞:读入源句的编码器,和写出译文的解码器,后者通过交叉注意力读取编码器的输出。它还在每次残差相加之后做归一化,用固定的正弦波表示位置,MLP 中用 ReLU。点击标签跳到对应的改动。

    6 层编码器 + 6 层解码器 · dmodel 512

  3. 翻译一个句子

    编码器把英文句子读一遍。然后解码器像 GPT-2 一样一次写一个德语词元,每一趟都读取编码器的输出。训练时直接喂入正确译文(教师强制),因果掩码遮住未来,所以所有位置一趟跑完。

    编码器一次 · 解码器每词元一次

  4. 三种注意力

    一共有三种注意力,它们都计算 softmax(Q·Kᵀ / √dk) · V。编码器能看到整个源句,解码器只能看到之前的目标词元,交叉注意力让每个目标词元都能看到整个源句。GPT-2 只有中间那种。

    源 × 源 · 目标 × 目标 · 目标 × 源

  5. 交叉注意力

    在交叉注意力中,查询来自解码器,键和值来自编码器的输出,所以分数矩阵是目标 × 源,不是方阵。每一行都在寻找它接下来需要的源词:‘gesehen’ 回头看向 ‘seen’,尽管德语把它挪到了句末。把鼠标停在格子上看看。

    Q 7 × dk · Kᵀ dk × 6 → 7 × 6

  6. 后置 LN vs 前置 LN

    2017 年的模型在每次残差相加之后归一化(后置 LN),这把 LayerNorm 放在了残差流的主路径上;深层的后置 LN 模型需要很长的学习率预热(论文中是 4,000 步)。GPT-2 则对每个子层读取的副本做归一化(前置 LN),所以主路径只做加法。

    LN(x + f(x)) → x + f(LN(x))

  7. 正弦位置编码

    注意力不管顺序,所以两个模型都给每个词元加上一个位置向量。2017 年的模型用正弦和余弦波算出它,没有参数,适用于任何位置;GPT-2 则学习一张 1,024 行的表。移动 k 个位置,就把每一对 sin/cos 旋转一个固定角度。

    PE(pos, 2i) = sin(pos / 10000^(2i/dmodel))

代码

# positions: fixed sine and cosine waves, added to embeddings scaled by √d_model
pe[:, 0::2] = torch.sin(pos * 10000 ** (-i2 / d_model))
pe[:, 1::2] = torch.cos(pos * 10000 ** (-i2 / d_model))
x = embed(src) * math.sqrt(d_model) + pe[:len(src)]
# encoder layer, 6 times: LayerNorm after each residual add (post-LN)
x = norm1(x + self_attn(x, x, x))                     # no mask
x = norm2(x + ffn(x))                                 # ReLU
# decoder layer, 6 times
y = norm1(y + self_attn(y, y, y, mask=causal))
y = norm2(y + cross_attn(q=y, k=memory, v=memory))    # memory = encoder output
y = norm3(y + ffn(y))
# GPT-2, for comparison: pre-LN
h = x + attn(ln_1(x)); out = h + mlp(ln_2(h))

延伸阅读