模型内部 · 注意力
- 图中数值
- 玩具规模
- 词元
- 5
- dmodel
- 8 GPT-2 768
- dhead
- 4 GPT-2 64
- 头
- 2 GPT-2 12
全部步骤
投影
X(经过 ln1 之后)分别乘以 WQ、WK 和 WV;GPT-2 还会加上偏置,玩具模型省略了。查询表示一个词元在找什么,键表示它能提供什么,值表示它被选中时交出什么。GPT-2 用一次 GEMM(X · Wqkv)同时算出三者。
GPT-2 [N×768]·[768×2304] + b · 每词元 3.5 MFLOPs分数
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]缩放
除以 √dhead。点积会随维度增大;缩放让 softmax 一开始就不至于饱和。
玩具 ÷ 2 · GPT-2 ÷ 8因果掩码
位置 i 预测词元 i + 1,所以它只能看 ≤ i 的位置:上三角被设为 −∞。训练时所有位置同时预测,如果没有掩码,每个位置都能直接读到下一个词。
因果:j > i → −∞Softmax
对每一行做 softmax,把分数变成和为 1 的权重;−∞ 经过 exp 变成 0。每一行就是一个词元的注意力分布。这些玩具权重是随机的,所以图案本身没有意义;“前向传播”一页展示了 GPT-2 真实的头。
A = softmax(S / √d + mask)加权求和
按注意力权重对 V 的各行加权求和。输出的第 i 行融合了所有可见词元的值,颜色也随之混合。右侧面板拆解当前这一行。
GPT-2 12 × [N×N]·[N×64]输出投影
把各个头的输出拼接起来,乘以 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