多头自注意力:QKV 分裂与因果掩码


文档摘要

多头自注意力:QKV 分裂与因果掩码 本节摘要:一次线性投影、三个视图、H 个并行头、一个掩码——这是模型实际用的注意力块。注意力是让一个 token 的表征从同序列其他 token 拉信息的函数;自注意力指 Q/K/V 都从同一输入派生;多头指投影拆成 H 个并行注意力问题,输出拼接再投影。本节建这个块的高效实现:一个 D→3D 的线性层,切成三视图,重塑成 H 个大小 D//H 的头;缩放点积、因果掩码、softmax、加权和全作批张量运算,头在加速器上并行。你会学到 QKV 分裂的等价性、头重塑、 缩放、因果掩码,以及逐头注意力权重的可解释性检查。 对应原课程:Phase 19 · Lesson 33 · (原英文 )。本节属「从零构建 GPT」赛道第四节。

多头自注意力:QKV 分裂与因果掩码

本节摘要:一次线性投影、三个视图、H 个并行头、一个掩码——这是模型实际用的注意力块。注意力是让一个 token 的表征从同序列其他 token 拉信息的函数;自注意力指 Q/K/V 都从同一输入派生;多头指投影拆成 H 个并行注意力问题,输出拼接再投影。本节建这个块的高效实现:一个 D→3D 的线性层,切成三视图,重塑成 H 个大小 D//H 的头;缩放点积、因果掩码、softmax、加权和全作批张量运算,头在加速器上并行。你会学到 QKV 分裂的等价性、头重塑、sqrt(d_head) 缩放、因果掩码,以及逐头注意力权重的可解释性检查。

对应原课程:Phase 19 · Lesson 33 · multihead-self-attention(原英文 phases/19-capstone-projects/33-multihead-self-attention/docs/en.md)。本节属「从零构建 GPT」赛道第四节。

学习目标

阅读完本节,你应当能够:

  1. 实现单线性层的批 Q/K/V 投影,拆成 H 个头。
  2. 计算带正确归一化与 dtype 处理的缩放点积注意力。
  3. 应用防止位置注意未来位置的因果掩码。
  4. 检查固定输入的逐头注意力权重,推理每个头看什么。
  5. 在玩具任务上训练小注意力块,看损失随头特化下降。

一、问题与直觉

高效实现模式是一个 D→3D 的线性层,切成三视图,重塑成 H 个大小 D//H 的头。matmul、softmax、加权和作批张量运算,头在加速器上并行。本节建这个块,加因果掩码,让同一代码可作 decoder-only 语言模型的注意力层。

形状契约:输入 (B, T, D)、输出 (B, T, D)、掩码 (T, T) 可广播;块内中间张量 (B, H, T, d_head),其中 d_head = D//H,约束 D % H == 0。两个线性层(QKV 投影与输出投影)是块内唯一参数;掩码、softmax、matmul、重塑全无参。

二、从零实现

QKV 分裂:朴素实现三独立线性层(各一),高效实现一层输出 3D 特征再切。两者数学等价——三个 (D, D) 权重的独立矩阵乘,正是由它们堆叠的 (3D, D) 权重的一次矩阵乘。高效版更快(加速器起一次 matmul 而非三次),也更易初始化(三子矩阵在同参数张量里一起初始化)。

头重塑:分裂后 Q/K/V 各 (B, T, D)。变成 H 个并行注意力问题,重塑成 (B, T, H, d_head) 再转置成 (B, H, T, d_head)——头维挨批维,PyTorch 把逐头注意力当跨 B*H 独立实例的批运算。d_head 维留末尾,score matmul Q @ K.transpose(-2, -1) 收缩它,结果 (B, H, T, T) 逐头注意力分。

缩放:softmax 前分数除 sqrt(d_head)。不缩放,点积随 d_head 长而长,把 softmax 推进「一项几乎全部质量、其余微乎其微」的区域,该区梯度极小、学习停滞。除 sqrt(d_head) 让分数方差跨头大小大致恒定。

因果掩码:decoder-only 语言模型预测下一 token 时只能条件于过去。掩码强制:softmax 前,(T, T) 分数矩阵对角线以上每项替成负无穷,softmax 后这些位置权重零。掩码构造时注册为 buffer(与模型同设备、非梯度图一部分),覆盖块将见的最大上下文长;前向切左上 (T, T) 角。

scores = Q @ K.transpose(-2, -1) / math.sqrt(d_head) # (B, H, T, T) scores = scores.masked_fill(self.mask[:T, :T] == 0, float("-inf")) weights = F.softmax(scores, dim=-1) context = weights @ V # (B, H, T, d_head)

三、输出投影与权重检查

逐头上下文向量 (B, H, T, d_head) 转置回 (B, T, H, d_head),重塑成 (B, T, D),应用最终 (D, D) 线性投影。输出投影让模型混合头——没它,H 头只能经后层重组,块被人为约束。

权重检查:本节前向暴露 return_weights=True 标志,设了就返 (B, H, T, T) 逐头权重。demo 打一头在短输入上的权重热图,你可看到因果三角结构与逐位置焦点。训练模型里不同头学不同模式:有的注意紧邻前 token,有的注意序列开头,有的几乎均匀铺开。检查钩子是可解释性工作的入口。

训练 demo:main.py 底部把注意力块接到微型 LM 头,在复制任务上训。输入每行是跨上下文复制的单个随机 id,目标是输入移位一——模型须学「下一 token 同前一」。损失交叉熵。H=4、D=32、T=12、词表 64,三 epoch 在 CPU 上损失从随机(log(64)≈4.16)降到远低于 1.0。重点不是训有用模型,是确认梯度流过块每片、头在答案明显的问题上学到东西。

四、可复用产物

main.py 定义 MultiHeadSelfAttention,持两线性层与注册掩码 buffer。前向:投影、重塑、打分、掩码、softmax、加权、重塑、再投影。demo 建小模型(token+位置嵌入+注意力+LM 头),在复制任务训三 epoch,打损失曲线与逐头注意力热图。code/tests/test_attention.py 钉形状契约、因果性、softmax 性质、头分裂性质、梯度流。把 n_heads 从 4 调 8(保 d_model=32d_head=4),看热图变化。

五、框架对比

本节实现是「原 Transformer 注意力 + 因果掩码」的 nanoGPT 风格高效实现。现代生产 transformer 加三样本节不做的事:feed-forward 块(注意力后接两层 MLP 加残差加 layer norm,RoPE 位置编码,以及推理用 KV 缓存。这三者分别属下一节(transformer 块)、位置编码扩展、推理节。本节块兼容 RoPE——只需在 matmul 前变换 Q 与 K。

六、练习

  1. 因果性:构造 (T, T) 输入,确认位置 t 的输出只依赖位置 ≤t 的输入(梯度或扰动测试)。
  2. 头特化:训 100 步,打每头权重热图,确认不同头学不同模式(前 token/序列开头/均匀)。
  3. 缩放影响:去掉 /sqrt(d_head),确认大 d_head 下学习停滞。
  4. 头数权衡:n_heads 4→8(保 D),看热图与损失曲线变化。
  5. QKV 等价性:分别用三独立线性与单 3D 线性,确认输出在同样初始化下数学等价。

本节要点回顾

  1. 单层 3D 投影:D→3D 切 Q/K/V,等价于三独立线性但更快。
  2. 头重塑:(B, T, D)(B, H, T, d_head),头维挨批维并行。
  3. sqrt(d_head) 缩放:防 softmax 进低梯度区,保分数方差跨头大小恒定。
  4. 因果掩码:对角线以上替负无穷,softmax 后权重零,buffer 注册非梯度图。
  5. 输出投影混合头:没它头只能经后层重组,块被人为约束。
  6. 逐头权重可解释:热图显因果三角与头特化,是可解释性入口。

下一节,我们把注意力块叠进「Transformer 块」——加 feed-forward、残差、层归一化。


发布者: 作者: Rohit Gupta 转发
评论区 (0)
U