模型内部 · 注意力

注意力

因果自注意力 · 第 1 块 · 头 1

图中数值
玩具规模
词元
5
dmodel
8 GPT-2 768
dhead
4 GPT-2 64
头
2 GPT-2 12

全部步骤

  1. 投影

    X(经过 ln1 之后)分别乘以 WQ、WK 和 WV;GPT-2 还会加上偏置,玩具模型省略了。查询表示一个词元在找什么,键表示它能提供什么,值表示它被选中时交出什么。GPT-2 用一次 GEMM(X · Wqkv)同时算出三者。

    GPT-2 [N×768]·[768×2304] + b · 每词元 3.5 MFLOPs

  2. 分数

    Q 乘以 K 的转置。第 i 行第 j 列是词元 i 的查询与词元 j 的键的点积:i 应该多关注 j。每个头只在自己那一段 dmodel / heads 个数上工作(GPT-2 中是 64 个,这里是 4 个)。这张表是 N × N 的:这里有 25 格,在 1,024 个词元时每个头超过一百万格,所以上下文翻倍,这一步的计算量就翻四倍。

    GPT-2 12 × [N×64]·[64×N]

  3. 缩放

    除以 √dhead。点积会随维度增大;缩放让 softmax 一开始就不至于饱和。

    玩具 ÷ 2 · GPT-2 ÷ 8

  4. 因果掩码

    位置 i 预测词元 i + 1,所以它只能看 ≤ i 的位置:上三角被设为 −∞。训练时所有位置同时预测,如果没有掩码,每个位置都能直接读到下一个词。

    因果:j > i → −∞

  5. Softmax

    对每一行做 softmax,把分数变成和为 1 的权重;−∞ 经过 exp 变成 0。每一行就是一个词元的注意力分布。这些玩具权重是随机的,所以图案本身没有意义;“前向传播”一页展示了 GPT-2 真实的头。

    A = softmax(S / √d + mask)

  6. 加权求和

    按注意力权重对 V 的各行加权求和。输出的第 i 行融合了所有可见词元的值,颜色也随之混合。右侧面板拆解当前这一行。

    GPT-2 12 × [N×N]·[N×64]

  7. 输出投影

    把各个头的输出拼接起来,乘以 WO 进行混合,再把结果加回残差流。几个小头可以各自追踪不同的关系,代价与一个大头相同。代码中:(B, T, 768) 在注意力中被拆成 (B, 12, T, 64),在这里再合并回去。

    GPT-2 [N×768]·[768×768] + b · 每词元 1.2 MFLOPs

代码

B, T, C = x.size()                                    # x = ln_1(h)
q, k, v = self.c_attn(x).split(self.n_embd, dim=2)    # one GEMM for all three
q = q.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)   # (B, 12, T, 64)
k = k.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
v = v.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
att = q @ k.transpose(-2, -1)                         # (B, 12, T, T)
att = att * (1.0 / math.sqrt(k.size(-1)))             # ÷ √64
att = att.masked_fill(self.bias[:, :, :T, :T] == 0, float("-inf"))
att = F.softmax(att, dim=-1)
y = att @ v                                           # (B, 12, T, 64)
y = y.transpose(1, 2).contiguous().view(B, T, C)      # heads side by side
y = self.c_proj(y)                                    # · W_O

延伸阅读