模型内部

前向传播

仅解码器 Transformer · GPT-2 small 真实运行

图中数值
真实 GPT-2
层
12
dmodel
768
头
12
词表
50,257
参数
124M

全部步骤

  1. 分词

    GPT-2 的字节级 BPE 把文本切成词元,每个词元是 50,257 项词表中的一个 id;Ġ 表示以空格开头的词元。点击词元可查看合并过程。

    ids [464, 3797, 3332, 319, 262]

  2. 嵌入

    每个 id 选出嵌入矩阵 WE 中对应的一行,再加上 WP 中对应位置的那一行;色条显示 768 个真实数值中的前 16 个。点击色条可查看查表过程。

    [5 × 768]

  3. 注意力

    每个位置从自身和之前的位置汇集信息:用它的查询(它在找什么)去和它们的键(它们能提供什么)打分,再按分数对它们的值加权混合;之后的词元被掩码遮住。把鼠标停在一条车道上查看它的权重,并在下方逐个查看全部 144 个真实的头。点击 attn 查看每一次乘法。

    12 个头 × [5 × 5]

  4. MLP

    MLP 单独处理每个位置:把 768 个数扩展到 3,072 个神经元,施加 GELU,再投影回去。ln1 和 ln2 这两块薄板在 attn 和 mlp 读取之前先把残差流归一化。点击 mlp 查看乘法过程。

    [5 × 768] → [5 × 3072]

  5. 第 2–12 块

    另外十一个块形状相同,但各有各的权重,它们把结果依次加到残差流上。右侧是 logit 透镜:如果模型在每一块之后就停下,最后一个位置会预测什么。车道颜色勾勒出信息的来源;色条显示各自的份额。

    12 个块约 8500 万参数 · 嵌入约 3900 万(WE,同时复用为 WU)

  6. 反嵌入

    只有最后一个位置做预测:lnf 把它归一化,反嵌入矩阵 WU(即 LM head,GPT-2 让它与 WEᵀ 共享权重)用一次点积为全部 50,257 个词元各打一个分。softmax 把分数变成这些概率,它们都来自 GPT-2 small 的真实运行。点击 lnf 或 WU 可逐步查看。

    [1 × 768] · WEᵀ → [1 × 50,257]

  7. 选出下一个词元

    本页总是取概率最高的词元(贪心解码),这样每个数都是真实的;对话模型则是采样,见“反嵌入与采样”。

    贪心 → Ġfloor · 7.6%

代码

idx = torch.tensor([enc.encode("The cat sat on the")])   # (1, T) token ids
x = wte(idx) + wpe(torch.arange(idx.size(1)))             # (1, T, 768)
for block in h:                                           # 12 blocks, own weights
    x = x + block.attn(block.ln_1(x))
    x = x + block.mlp(block.ln_2(x))
logits = lm_head(ln_f(x)[:, -1, :])                       # last position only
idx_next = logits.argmax(-1, keepdim=True)                # greedy
idx = torch.cat((idx, idx_next), dim=1)                   # append, next pass

延伸阅读