观文听傑

返回

上一篇的加性注意力让第 tt 个解码步直接读取全部源状态,突破了定长上下文瓶颈。但循环编码器仍要先完成第 1、2、3……步,循环解码器也要逐 token 生成;序列越长,时间轴上的串行依赖越明显。

Transformer 的核心转向是自注意力(Self-Attention):让同一层中的每个位置用查询(Query)寻找其他位置的键(Key),再汇总对应的值(Value)。本文只讲透缩放点积自注意力(Scaled Dot-Product Attention)这一个算子,并把与正确使用它不可分的多头拆分、位置信息和掩码一起说明。前馈网络、完整 Encoder–Decoder 堆叠和大模型训练留到后续文章。

01 循环结构的限制不只是“记不住”#

循环网络的状态更新是:

ht=f(xt,ht1)h_t=f(x_t,h_{t-1})

即使 LSTM 缓解了长程梯度问题,即使加性注意力能回看全部源状态,计算 hth_t 仍必须等待 ht1h_{t-1}。如果第 1 个词要与第 100 个词交换信息,它至少要跨越许多递推或等到解码查询发生。

自注意力把一层序列写成矩阵:

XRN×L×DmodelX\in\mathbb{R}^{N\times L\times D_{model}}

同一层一次生成所有位置的 Q,K,VQ,K,V,再用一个 [L,L] 关系矩阵完成位置间的信息交换:

循环层:
x1 ─► h1 ─► h2 ─► h3 ─► h4          时间轴串行
          x2    x3    x4

自注意力层:
x1 ─┬────────────► every output position
x2 ─┼────────────► 通过同一次 QK^T 建立 L×L 连接
x3 ─┼────────────►
x4 ─┴────────────►
text

这让训练阶段的序列位置更易并行,但关系矩阵的时间和显存通常随 L2L^2 增长。并行不是免费消除复杂度,而是把串行递推换成密集矩阵计算。

02 Query、Key、Value 各自做什么?#

对输入位置表示 XX 做三组线性投影:

Q=XWQ,K=XWK,V=XWVQ=XW^Q,\qquad K=XW^K,\qquad V=XW^V

单头时可设:

Q,KRN×L×dk,VRN×L×dvQ,K\in\mathbb{R}^{N\times L\times d_k}, \qquad V\in\mathbb{R}^{N\times L\times d_v}

ii 个位置的查询 qiq_i 与第 jj 个位置的键 kjk_j 做点积,决定 iijj 读取多少;真正被加权汇总的是 vjv_j

Attention(Q,K,V)=softmax(QKdk+M)V\operatorname{Attention}(Q,K,V) =\operatorname{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}+M\right)V
X [N,L,D_model]
 ├─► Linear Q ─► Q [N,L,d_k] ─────────┐
 ├─► Linear K ─► K [N,L,d_k] ─► 转置 ├─► QK^T / sqrt(d_k) [N,L,L]
 └─► Linear V ─► V [N,L,d_v]          │                 │
                                      mask M ────────────┤

                                               softmax(dim=key_position)
                                                        │ A [N,L,L]

                                              A @ V = O [N,L,d_v]
text

“查询、键、值”是计算角色,不是三份独立输入数据。在自注意力中它们都由同一个 XX 投影得到;在交叉注意力中,查询可来自解码器,键和值来自编码器。

03 为什么点积要除以 dk\sqrt{d_k}#

假设 qqkk 各维独立、均值 0、方差 1,则点积:

qk=r=1dkqrkrq^\top k=\sum_{r=1}^{d_k}q_rk_r

其方差约为 dkd_k,标准差约为 dk\sqrt{d_k}。维度增大时,未经缩放的 logits 会越来越极端,softmax 接近 one-hot,非最大项梯度变小。

例如两个键的未缩放分数为 [8,0]

softmax([8,0])[0.9997,0.0003]\operatorname{softmax}([8,0])\approx[0.9997,0.0003]

dk=16d_k=16,缩放后是 [2,0]

softmax([2,0])[0.8808,0.1192]\operatorname{softmax}([2,0])\approx[0.8808,0.1192]

缩放不是为了让权重“平均”,而是让点积尺度在不同头维度下更可控,使训练初期不过早饱和。若框架 API 已经完成缩放,再手动除一次会把分布错误地变平。

04 用两个 token 手算完整前向#

令单个样本、单头、L=2,dk=dv=2L=2,d_k=d_v=2

Q=K=[1001],V=[2004]Q=K=\begin{bmatrix}1&0\\0&1\end{bmatrix}, \qquad V=\begin{bmatrix}2&0\\0&4\end{bmatrix}

缩放前的相似度矩阵是单位矩阵,除以 2\sqrt2 后:

S=QK2=[0.707000.707]S=\frac{QK^\top}{\sqrt2} =\begin{bmatrix}0.707&0\\0&0.707\end{bmatrix}

逐行 softmax,令 a=e0.707e0.707+10.6698a=\frac{e^{0.707}}{e^{0.707}+1}\approx0.6698

A=[a1a1aa][0.66980.33020.33020.6698]A=\begin{bmatrix}a&1-a\\1-a&a\end{bmatrix} \approx \begin{bmatrix}0.6698&0.3302\\0.3302&0.6698\end{bmatrix}

输出为:

O=AV[1.33961.32080.66042.6792]O=AV \approx \begin{bmatrix} 1.3396&1.3208\\ 0.6604&2.6792 \end{bmatrix}

第一位置仍更重视自己,却读取了第二位置约 33% 的值;第二位置同理。注意输出不是对 QQKK 加权,而是对 VV 加权。

若这是自回归语言模型,第一个位置不能偷看第二个 token。加入因果掩码后,第一行第二列变为 -\infty

Scausal=[0.70700.707]S_{causal}=\begin{bmatrix}0.707&-\infty\\0&0.707\end{bmatrix}

第一行权重变为 [1,0],输出严格等于第一个 value;第二位置仍能读取位置 1 和 2。

05 多头不是“把同一注意力复制几遍”#

多头注意力(Multi-Head Attention)把 DmodelD_{model} 投影成 HH 个子空间。常见设置为 dk=dv=Dmodel/Hd_k=d_v=D_{model}/H

headr=Attention(QWrQ,KWrK,VWrV)\operatorname{head}_r =\operatorname{Attention}(QW_r^Q,KW_r^K,VW_r^V) MultiHead(Q,K,V)=Concat(head1,,headH)WO\operatorname{MultiHead}(Q,K,V) =\operatorname{Concat}(\operatorname{head}_1,\ldots,\operatorname{head}_H)W^O
张量形状
输入 x[N,L,D_model]
投影后 q,k,v[N,L,H*d_head]
拆头并换轴[N,H,L,d_head]
每头权重[N,H,L,L]
每头输出[N,H,L,d_head]
拼接后[N,L,H*d_head]
输出投影[N,L,D_model]

不同头有独立投影,因而可以学习不同关系子空间;但不能仅凭某个头的图案给它命名为“句法头”或“实体头”。头也可能冗余、塌缩或在不同随机种子下交换角色。

D_model 必须能按头数切分,是最基本的形状约束:

assert model_dim % num_heads == 0
head_dim = model_dim // num_heads
python

06 自注意力为什么仍需要位置信息?#

如果没有任何位置表示,自注意力对 token 排列是置换等变的:把输入行按同一种顺序重排,输出也只会跟着重排。模型知道「狗」「咬」「人」有哪些内容,却没有天然坐标区分「狗咬人」和「人咬狗」。

原始 Transformer 把正弦位置编码(Sinusoidal Positional Encoding)加到 token embedding:

PE(pos,2i)=sin(pos/100002i/Dmodel)PE(pos,2i)=\sin\left(pos/10000^{2i/D_{model}}\right) PE(pos,2i+1)=cos(pos/100002i/Dmodel)PE(pos,2i+1)=\cos\left(pos/10000^{2i/D_{model}}\right)
token_ids [N,L] ─► Embedding [N,L,D]
position  [L]   ─► Position  [L,D]
                         │ broadcast batch

                    X = token + position [N,L,D]

                    self-attention
text

学习式绝对位置、相对位置偏置、旋转位置编码等都在改变“位置如何进入相似度或表示”,但都不应与 padding mask 混为一谈:位置编码提供顺序;mask 禁止某些连接。

07 padding mask 与 causal mask 阻止不同错误#

两类 mask 经常同时出现:

  • 键 padding mask:每个样本真实长度不同,任何查询都不应读取补齐键位置。典型语义为 [N,L]
  • 因果 mask(Causal Mask):语言模型位置 ii 不应读取 j>ij>i 的未来内容。典型语义为 [L,L] 下三角。

mask 通常屏蔽的是 key 列。一个 padding 查询行仍可能产生输出;若后续不需要它,应在残差输出或目标损失处再次屏蔽。只屏蔽查询行而保留 padding 键,会让真实 token 读取无效内容。

08 不调用封装,写出可检查的多头本体#

下面用 valid_keys=True 表示真实 token,内部统一把不可见位置填为 -\infty

transpose 后调用 contiguous()view,避免把非连续内存按错误步长解释。也可以使用 reshape,但仍应理解最终布局从 [N,H,L,d] 变为 [N,L,H,d] 后才能拼头。

训练时注意力通常位于残差块中:

x ─► LayerNorm ─► SelfAttention ─► dropout ─► + ─► y
└──────────────────────────────────────────────▲
text

只输出注意力结果而不加残差、归一化和后续前馈层,不等于一个完整 Transformer block;本文刻意只验证注意力算子。

09 与 PyTorch 2.13 当前 SDPA API 对齐#

PyTorch 2.13 的 torch.nn.functional.scaled_dot_product_attention 接收形如 [N,H,L,d] 的 Query 和 [N,H,S,d] 的 Key/Value,并在可能时选择优化内核:

当前官方契约中有四个高风险点:

  1. 默认缩放已经是 1/dk1/\sqrt{d_k},不要在传入 q 前再除一次。
  2. 布尔 attn_mask=True 表示该连接允许参与;浮点 mask 则直接加到分数。
  3. 当前接口不允许同时显式传 attn_maskis_causal=True;padding 与 causal 需要按 API/版本选择合并策略或更高层封装。
  4. 该函数只要 dropout_p>0 就会应用 dropout,不会自动读取外层模块的 training;模块中应传 self.p if self.training else 0.0

该函数当前仍标为 Beta。固定 PyTorch 版本、对 mask 做数值单测,比只依赖“代码能运行”更可靠。

nn.MultiheadAttention(batch_first=True) 则接收 [N,L,D]。其布尔 key_padding_mask [N,S]True 表示忽略,与上面的 SDPA 布尔 mask 相反:

mha = nn.MultiheadAttention(
    embed_dim=32, num_heads=4, dropout=0.1, batch_first=True
)
x = torch.randn(2, 6, 32)
key_padding_mask = ~valid_keys  # MHA 中 True = 忽略

mha.train()
output, per_head = mha(
    x, x, x,
    key_padding_mask=key_padding_mask,
    need_weights=True,
    average_attn_weights=False,
)  # output [2,6,32], per_head [2,4,6,6]
python

生产路径若不需要权重,设置 need_weights=False 更容易使用优化后的 scaled dot-product attention;诊断时再在小 batch 上取每头权重。

10 一个语言模型训练步的输入输出#

自回归语言模型把 token 序列错开一位:

原序列: <bos> 今 天 下 雨 <eos> <pad>
输入 x: <bos> 今 天 下 雨
标签 y:  今   天 下 雨 <eos>
text

输入先加位置表示,经过带 causal mask 的若干 Transformer block,再投影到词表 logits:

tokens [N,L]X[N,L,D]Z[N,L,D]logits [N,L,V]\text{tokens }[N,L] \rightarrow X[N,L,D] \rightarrow Z[N,L,D] \rightarrow \text{logits }[N,L,V]
logits = model(input_ids, valid_keys=input_ids.ne(pad_id), causal=True)
# logits [N,L,V],未经 softmax
loss = nn.functional.cross_entropy(
    logits.reshape(-1, logits.shape[-1]),
    target_ids.reshape(-1),
    ignore_index=pad_id,
)

optimizer.zero_grad(set_to_none=True)
loss.backward()
grad_norm = nn.utils.clip_grad_norm_(
    model.parameters(), max_norm=1.0, error_if_nonfinite=True
)
optimizer.step()
python

ignore_index 只使目标 padding 不贡献损失;它不会阻止真实查询在注意力中读取 padding 键。反过来,注意力 key mask 也不会自动让 padding 目标不计损失。二者必须分别存在。

训练时整段目标已知,可用 causal mask 并行计算所有位置;推理时未来 token 尚不存在,仍然要自回归逐步生成。键值缓存(KV Cache)可以复用已生成 token 的 Key/Value,避免每步重新计算整段历史,但不会让依赖关系消失。

11 一条可执行的调试路径#

  1. 先过拟合一个 batch。 用 2 条、长度 4 的复制或 next-token 数据,确认损失可接近 0。
  2. 检查权重和。 dropout 关闭时,weights.sum(-1) 应约为 1;padding 列应为 0。
  3. 做未来泄漏测试。 固定前缀,只替换位置 tt 之后的 token;causal 模型在 t\le t 位置的 logits 应完全不变。
  4. 做排列测试。 暂时移除位置表示,成对重排输入,输出应同样重排;加入位置后该对称性应被打破。
  5. 比较手写与官方输出。 关闭 dropout,复制投影参数或直接给相同 Q,K,VQ,K,V,使用 torch.testing.assert_close
  6. 记录注意力熵与 logits 范数。 全头长期均匀可能未学到关系;极早 one-hot 可能是缩放、初始化或 mask 问题。
  7. 用 profiler 看长度曲线。LL 翻倍,分别记录注意力矩阵显存、吞吐和数据加载时间,确认瓶颈是否真在 L2L^2 算子。

未来泄漏的最小测试尤其重要:模型的训练损失会因偷看答案而异常漂亮,普通形状断言却抓不到它。

12 最常见的“形状正确,语义错误”#

  • softmax 沿查询轴。 每个查询应沿 key 位置归一化,即最后一维。
  • 忘记转置 Key。 需要 [N,H,L,d] @ [N,H,d,S] 才得到 [N,H,L,S]
  • DmodelD_{model} 而不是 dkd_k 缩放每个头。 缩放由单头 Query/Key 宽度决定。
  • 手动缩放后又调用 SDPA。 重复除以平方根会让注意力过平。
  • 混淆 SDPA 与 MHA 的布尔 mask。 同一个 True 在两个接口里可表示相反语义。
  • 只 mask padding 查询,不 mask padding 键。 真实 token 仍可能读取补齐位置。
  • causal 三角方向反了。 打印一个 4×44\times4 mask,并用“替换未来不改变过去 logits”测试。
  • 所有键都被屏蔽。 对全 -\infty 行做 softmax 会产生非有限结果;入口拒绝空序列。
  • 没有位置表示却期待词序。 内容相同的排列无法仅靠无位置自注意力区分。
  • 拆头后直接 view 先换回 [N,L,H,d] 并确保内存布局正确。
  • 验证时 SDPA 仍传训练 dropout。 当前函数会按 dropout_p 无条件应用 dropout。
  • 只看平均头权重。 平均会抹去头间差异;诊断时取 [N,H,L,S],生产时可关闭权重返回。
  • 把训练并行误解为生成并行。 causal 训练能一次计算所有已知标签,开放式推理仍依赖已生成前缀。

13 它会在哪些场景失败?#

  • 超长序列。 标准注意力构造 L×LL\times L 分数,长文档、视频和高分辨率网格的显存迅速增长。
  • 局部模式占主导。 没有适当位置归纳偏置时,小数据上可能不如卷积或精心设计的局部模型。
  • 精确外推到更长长度。 训练长度、位置表示和数值范围都可能限制长度外推。
  • 因果生成延迟。 训练能并行位置,逐 token 推理仍受串行采样和 KV cache 带宽约束。
  • 注意力不是可靠检索。 有限精度的加权平均可能混合多个相似值,不能替代带标识符的精确数据库读取。
  • 权重不等于解释。 改变 Value 或后续层可能在权重相似时改变答案;解释需要干预和多种证据。
  • 数据捷径。 全局连接让模型更容易利用非因果元数据、模板位置或重复样本,数据切分仍是第一道防线。

稀疏注意力、线性注意力、局部窗口和 FlashAttention 分别优化连接模式、数学近似或内存访问;不能仅因都“更快”就认为它们等价。

14 与前后方法怎样区分?#

方法位置间路径训练时序列并行主要代价/限制
RNN/LSTM逐步状态递推长路径、吞吐受串行依赖限制
循环 + 加性注意力解码查询读取全部编码状态编码/解码仍递推每步对源序列打分
Transformer 自注意力一层内所有位置直接连接标准形式为 O(L2)O(L^2) 关系矩阵
卷积序列模型固定局部窗口逐层扩大长程关系需更多层或膨胀卷积

自注意力算子本身没有定义完整 Transformer。标准块还包含残差连接、LayerNorm、逐位置前馈网络、dropout;Encoder–Decoder Transformer 还包含跨源—目标的交叉注意力。先把一个算子的轴、缩放和 mask 测对,再讨论堆叠深度和架构变体。

15 今天真正需要记住什么?#

  1. 自注意力用 QKQK^\top 建立“哪个查询读取哪个键”的关系,再用归一化权重汇总 VV;输出位置可以在同一层直接交换信息。
  2. 除以 dk\sqrt{d_k} 是为控制点积方差和 softmax 饱和,单头缩放宽度是 dkd_k
  3. 多头通过独立投影学习不同子空间,数据流核心是 [N,L,D] → [N,H,L,d] → [N,H,L,L] → [N,L,D]
  4. 位置表示提供顺序,padding mask 禁止读取补齐键,causal mask 禁止读取未来;三者职责不同。
  5. PyTorch 2.13 的 SDPA 与 MultiheadAttention 对布尔 mask 的 True 语义不同,且 SDPA dropout 不自动随 eval() 关闭,必须写契约测试。

16 思考题与小练习#

  1. 延续二 token 手算,把 VV 改为 [[1,1],[3,-1]],计算无 mask 与 causal mask 下两个位置的输出;说明 mask 改变的是权重可见性而不是 Value 本身。
  2. TransparentSelfAttention 写未来泄漏测试:随机替换位置 3 之后的输入,验证 causal 模式下位置 0–3 的输出不变;再故意把三角 mask 翻转,观察测试如何失败。
  3. 固定 D_model=128,比较 H∈{1,4,8} 的每头宽度、注意力矩阵元素数、参数量和吞吐。解释增加头数为何不必然增加 QKV 投影参数,却会改变每头的缩放与表示子空间。

相关工作#

17 下一篇预告#

自注意力已经让位置在同一层直接交换信息,但如果没有残差、归一化和逐位置非线性,整层仍只是一次加权混合。下一篇将把这些部件组装成 Transformer block,比较 Pre-LN 与 Post-LN 的数据流和梯度路径,并追踪一个 token 如何经过注意力子层与前馈子层完成更新。

所有词怎样一次看见彼此?Transformer 的缩放点积自注意力与掩码
https://zwjcode.cn/blog/transformer-scaled-dot-product-self-attention-masks
作者
发布于 2026年9月6日
版权协议 CC BY-NC-SA 4.0
评论加载似乎遇到了问题,请尝试刷新页面。