所有词怎样一次看见彼此?Transformer 的缩放点积自注意力与掩码
从循环注意力的串行瓶颈出发,手算 Query、Key、Value 的缩放点积,解释多头分工、位置编码与两类掩码,并对齐 PyTorch 2.13 当前 API。
上一篇的加性注意力让第 个解码步直接读取全部源状态,突破了定长上下文瓶颈。但循环编码器仍要先完成第 1、2、3……步,循环解码器也要逐 token 生成;序列越长,时间轴上的串行依赖越明显。
Transformer 的核心转向是自注意力(Self-Attention):让同一层中的每个位置用查询(Query)寻找其他位置的键(Key),再汇总对应的值(Value)。本文只讲透缩放点积自注意力(Scaled Dot-Product Attention)这一个算子,并把与正确使用它不可分的多头拆分、位置信息和掩码一起说明。前馈网络、完整 Encoder–Decoder 堆叠和大模型训练留到后续文章。
01 循环结构的限制不只是“记不住”#
循环网络的状态更新是:
即使 LSTM 缓解了长程梯度问题,即使加性注意力能回看全部源状态,计算 仍必须等待 。如果第 1 个词要与第 100 个词交换信息,它至少要跨越许多递推或等到解码查询发生。
自注意力把一层序列写成矩阵:
同一层一次生成所有位置的 ,再用一个 [L,L] 关系矩阵完成位置间的信息交换:
循环层:
x1 ─► h1 ─► h2 ─► h3 ─► h4 时间轴串行
x2 x3 x4
自注意力层:
x1 ─┬────────────► every output position
x2 ─┼────────────► 通过同一次 QK^T 建立 L×L 连接
x3 ─┼────────────►
x4 ─┴────────────►text这让训练阶段的序列位置更易并行,但关系矩阵的时间和显存通常随 增长。并行不是免费消除复杂度,而是把串行递推换成密集矩阵计算。
02 Query、Key、Value 各自做什么?#
对输入位置表示 做三组线性投影:
单头时可设:
第 个位置的查询 与第 个位置的键 做点积,决定 从 读取多少;真正被加权汇总的是 :
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“查询、键、值”是计算角色,不是三份独立输入数据。在自注意力中它们都由同一个 投影得到;在交叉注意力中,查询可来自解码器,键和值来自编码器。
03 为什么点积要除以 ?#
假设 和 各维独立、均值 0、方差 1,则点积:
其方差约为 ,标准差约为 。维度增大时,未经缩放的 logits 会越来越极端,softmax 接近 one-hot,非最大项梯度变小。
例如两个键的未缩放分数为 [8,0]:
若 ,缩放后是 [2,0]:
缩放不是为了让权重“平均”,而是让点积尺度在不同头维度下更可控,使训练初期不过早饱和。若框架 API 已经完成缩放,再手动除一次会把分布错误地变平。
04 用两个 token 手算完整前向#
令单个样本、单头、:
缩放前的相似度矩阵是单位矩阵,除以 后:
逐行 softmax,令 :
输出为:
第一位置仍更重视自己,却读取了第二位置约 33% 的值;第二位置同理。注意输出不是对 或 加权,而是对 加权。
若这是自回归语言模型,第一个位置不能偷看第二个 token。加入因果掩码后,第一行第二列变为 :
第一行权重变为 [1,0],输出严格等于第一个 value;第二位置仍能读取位置 1 和 2。
05 多头不是“把同一注意力复制几遍”#
多头注意力(Multi-Head Attention)把 投影成 个子空间。常见设置为 :
| 张量 | 形状 |
|---|---|
输入 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_headspython06 自注意力为什么仍需要位置信息?#
如果没有任何位置表示,自注意力对 token 排列是置换等变的:把输入行按同一种顺序重排,输出也只会跟着重排。模型知道「狗」「咬」「人」有哪些内容,却没有天然坐标区分「狗咬人」和「人咬狗」。
原始 Transformer 把正弦位置编码(Sinusoidal Positional Encoding)加到 token embedding:
token_ids [N,L] ─► Embedding [N,L,D]
position [L] ─► Position [L,D]
│ broadcast batch
▼
X = token + position [N,L,D]
│
self-attentiontext学习式绝对位置、相对位置偏置、旋转位置编码等都在改变“位置如何进入相似度或表示”,但都不应与 padding mask 混为一谈:位置编码提供顺序;mask 禁止某些连接。
07 padding mask 与 causal mask 阻止不同错误#
两类 mask 经常同时出现:
- 键 padding mask:每个样本真实长度不同,任何查询都不应读取补齐键位置。典型语义为
[N,L]。 - 因果 mask(Causal Mask):语言模型位置 不应读取 的未来内容。典型语义为
[L,L]下三角。
序列: [A, B, C, PAD]
key padding mask(列规则,每一行都屏蔽 PAD):
A B C P
A ✓ ✓ ✓ ×
B ✓ ✓ ✓ ×
C ✓ ✓ ✓ ×
P ✓ ✓ ✓ × ← 查询 PAD 是否保留还需在输出/损失端处理
causal mask(时间规则):
A B C P
A ✓ × × ×
B ✓ ✓ × ×
C ✓ ✓ ✓ ×
P ✓ ✓ ✓ ✓
组合后:同时满足“非未来”与“非 padding”才可见。textmask 通常屏蔽的是 key 列。一个 padding 查询行仍可能产生输出;若后续不需要它,应在残差输出或目标损失处再次屏蔽。只屏蔽查询行而保留 padding 键,会让真实 token 读取无效内容。
08 不调用封装,写出可检查的多头本体#
下面用 valid_keys=True 表示真实 token,内部统一把不可见位置填为 :
import math
import torch
from torch import nn
class TransparentSelfAttention(nn.Module):
def __init__(self, model_dim: int, num_heads: int, dropout: float = 0.0) -> None:
super().__init__()
assert model_dim % num_heads == 0
self.model_dim = model_dim
self.num_heads = num_heads
self.head_dim = model_dim // num_heads
self.qkv = nn.Linear(model_dim, 3 * model_dim)
self.output = nn.Linear(model_dim, model_dim)
self.dropout = nn.Dropout(dropout)
def _split_heads(self, x: torch.Tensor) -> torch.Tensor:
n, length, _ = x.shape
return x.view(n, length, self.num_heads, self.head_dim).transpose(1, 2)
# [N,H,L,d]
def forward(
self,
x: torch.Tensor, # [N,L,D]
valid_keys: torch.Tensor, # [N,L], True = 可读
causal: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]:
n, length, width = x.shape
assert width == self.model_dim
assert valid_keys.shape == (n, length)
assert valid_keys.dtype == torch.bool
assert valid_keys.any(dim=1).all()
q_raw, k_raw, v_raw = self.qkv(x).chunk(3, dim=-1)
q = self._split_heads(q_raw) # [N,H,L,d]
k = self._split_heads(k_raw) # [N,H,L,d]
v = self._split_heads(v_raw) # [N,H,L,d]
scores = q @ k.transpose(-2, -1) / math.sqrt(self.head_dim)
# [N,H,query=L,key=L]
allowed = valid_keys[:, None, None, :] # [N,1,1,L]
if causal:
lower_triangle = torch.ones(
length, length, dtype=torch.bool, device=x.device
).tril()
allowed = allowed & lower_triangle[None, None, :, :]
scores = scores.masked_fill(~allowed, float("-inf"))
weights = scores.softmax(dim=-1) # 沿 key 位置归一化
weights = self.dropout(weights)
attended = weights @ v # [N,H,L,d]
merged = attended.transpose(1, 2).contiguous().view(n, length, width)
return self.output(merged), weights # [N,L,D], [N,H,L,L]pythontranspose 后调用 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,并在可能时选择优化内核:
import torch
from torch.nn import functional as F
q = torch.randn(2, 4, 6, 8) # [N=2,H=4,L=6,d=8]
k = torch.randn(2, 4, 6, 8)
v = torch.randn(2, 4, 6, 8)
valid_keys = torch.tensor([
[True, True, True, True, False, False],
[True, True, True, True, True, True],
]) # [N,S]
# SDPA 布尔 attn_mask 中 True = 允许参与;广播到 [N,H,L,S]
allowed = valid_keys[:, None, None, :]
output = F.scaled_dot_product_attention(
q, k, v,
attn_mask=allowed,
dropout_p=0.0,
is_causal=False,
) # [2,4,6,8]python当前官方契约中有四个高风险点:
- 默认缩放已经是 ,不要在传入
q前再除一次。 - 布尔
attn_mask=True表示该连接允许参与;浮点 mask 则直接加到分数。 - 当前接口不允许同时显式传
attn_mask和is_causal=True;padding 与 causal 需要按 API/版本选择合并策略或更高层封装。 - 该函数只要
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:
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()pythonignore_index 只使目标 padding 不贡献损失;它不会阻止真实查询在注意力中读取 padding 键。反过来,注意力 key mask 也不会自动让 padding 目标不计损失。二者必须分别存在。
训练时整段目标已知,可用 causal mask 并行计算所有位置;推理时未来 token 尚不存在,仍然要自回归逐步生成。键值缓存(KV Cache)可以复用已生成 token 的 Key/Value,避免每步重新计算整段历史,但不会让依赖关系消失。
11 一条可执行的调试路径#
- 先过拟合一个 batch。 用 2 条、长度 4 的复制或 next-token 数据,确认损失可接近 0。
- 检查权重和。 dropout 关闭时,
weights.sum(-1)应约为 1;padding 列应为 0。 - 做未来泄漏测试。 固定前缀,只替换位置 之后的 token;causal 模型在 位置的 logits 应完全不变。
- 做排列测试。 暂时移除位置表示,成对重排输入,输出应同样重排;加入位置后该对称性应被打破。
- 比较手写与官方输出。 关闭 dropout,复制投影参数或直接给相同 ,使用
torch.testing.assert_close。 - 记录注意力熵与 logits 范数。 全头长期均匀可能未学到关系;极早 one-hot 可能是缩放、初始化或 mask 问题。
- 用 profiler 看长度曲线。 将 翻倍,分别记录注意力矩阵显存、吞吐和数据加载时间,确认瓶颈是否真在 算子。
未来泄漏的最小测试尤其重要:模型的训练损失会因偷看答案而异常漂亮,普通形状断言却抓不到它。
12 最常见的“形状正确,语义错误”#
- softmax 沿查询轴。 每个查询应沿 key 位置归一化,即最后一维。
- 忘记转置 Key。 需要
[N,H,L,d] @ [N,H,d,S]才得到[N,H,L,S]。 - 用 而不是 缩放每个头。 缩放由单头 Query/Key 宽度决定。
- 手动缩放后又调用 SDPA。 重复除以平方根会让注意力过平。
- 混淆 SDPA 与 MHA 的布尔 mask。 同一个
True在两个接口里可表示相反语义。 - 只 mask padding 查询,不 mask padding 键。 真实 token 仍可能读取补齐位置。
- causal 三角方向反了。 打印一个 mask,并用“替换未来不改变过去 logits”测试。
- 所有键都被屏蔽。 对全 行做 softmax 会产生非有限结果;入口拒绝空序列。
- 没有位置表示却期待词序。 内容相同的排列无法仅靠无位置自注意力区分。
- 拆头后直接
view。 先换回[N,L,H,d]并确保内存布局正确。 - 验证时 SDPA 仍传训练 dropout。 当前函数会按
dropout_p无条件应用 dropout。 - 只看平均头权重。 平均会抹去头间差异;诊断时取
[N,H,L,S],生产时可关闭权重返回。 - 把训练并行误解为生成并行。 causal 训练能一次计算所有已知标签,开放式推理仍依赖已生成前缀。
13 它会在哪些场景失败?#
- 超长序列。 标准注意力构造 分数,长文档、视频和高分辨率网格的显存迅速增长。
- 局部模式占主导。 没有适当位置归纳偏置时,小数据上可能不如卷积或精心设计的局部模型。
- 精确外推到更长长度。 训练长度、位置表示和数值范围都可能限制长度外推。
- 因果生成延迟。 训练能并行位置,逐 token 推理仍受串行采样和 KV cache 带宽约束。
- 注意力不是可靠检索。 有限精度的加权平均可能混合多个相似值,不能替代带标识符的精确数据库读取。
- 权重不等于解释。 改变 Value 或后续层可能在权重相似时改变答案;解释需要干预和多种证据。
- 数据捷径。 全局连接让模型更容易利用非因果元数据、模板位置或重复样本,数据切分仍是第一道防线。
稀疏注意力、线性注意力、局部窗口和 FlashAttention 分别优化连接模式、数学近似或内存访问;不能仅因都“更快”就认为它们等价。
14 与前后方法怎样区分?#
| 方法 | 位置间路径 | 训练时序列并行 | 主要代价/限制 |
|---|---|---|---|
| RNN/LSTM | 逐步状态递推 | 否 | 长路径、吞吐受串行依赖限制 |
| 循环 + 加性注意力 | 解码查询读取全部编码状态 | 编码/解码仍递推 | 每步对源序列打分 |
| Transformer 自注意力 | 一层内所有位置直接连接 | 是 | 标准形式为 关系矩阵 |
| 卷积序列模型 | 固定局部窗口逐层扩大 | 是 | 长程关系需更多层或膨胀卷积 |
自注意力算子本身没有定义完整 Transformer。标准块还包含残差连接、LayerNorm、逐位置前馈网络、dropout;Encoder–Decoder Transformer 还包含跨源—目标的交叉注意力。先把一个算子的轴、缩放和 mask 测对,再讨论堆叠深度和架构变体。
15 今天真正需要记住什么?#
- 自注意力用 建立“哪个查询读取哪个键”的关系,再用归一化权重汇总 ;输出位置可以在同一层直接交换信息。
- 除以 是为控制点积方差和 softmax 饱和,单头缩放宽度是 。
- 多头通过独立投影学习不同子空间,数据流核心是
[N,L,D] → [N,H,L,d] → [N,H,L,L] → [N,L,D]。 - 位置表示提供顺序,padding mask 禁止读取补齐键,causal mask 禁止读取未来;三者职责不同。
- PyTorch 2.13 的 SDPA 与 MultiheadAttention 对布尔 mask 的
True语义不同,且 SDPA dropout 不自动随eval()关闭,必须写契约测试。
16 思考题与小练习#
- 延续二 token 手算,把 改为
[[1,1],[3,-1]],计算无 mask 与 causal mask 下两个位置的输出;说明 mask 改变的是权重可见性而不是 Value 本身。 - 为
TransparentSelfAttention写未来泄漏测试:随机替换位置 3 之后的输入,验证 causal 模式下位置 0–3 的输出不变;再故意把三角 mask 翻转,观察测试如何失败。 - 固定
D_model=128,比较H∈{1,4,8}的每头宽度、注意力矩阵元素数、参数量和吞吐。解释增加头数为何不必然增加 QKV 投影参数,却会改变每头的缩放与表示子空间。
相关工作#
- Vaswani et al. (2017), Attention Is All You Need ↗:提出以多头自注意力为核心、无需循环与卷积的 Transformer。
- Shaw, Uszkoreit & Vaswani (2018), Self-Attention with Relative Position Representations ↗:将相对位置信息直接加入自注意力关系计算。
- Dai et al. (2019), Transformer-XL ↗:用片段级状态复用与相对位置缓解固定上下文和长度依赖。
- Dao et al. (2022), FlashAttention ↗:通过 IO 感知的精确注意力算法减少显存访问,而非改变注意力数学结果。
- Press, Smith & Lewis (2022), Train Short, Test Long: Attention with Linear Biases ↗:用注意力线性偏置研究长度外推与位置表示。
17 下一篇预告#
自注意力已经让位置在同一层直接交换信息,但如果没有残差、归一化和逐位置非线性,整层仍只是一次加权混合。下一篇将把这些部件组装成 Transformer block,比较 Pre-LN 与 Post-LN 的数据流和梯度路径,并追踪一个 token 如何经过注意力子层与前馈子层完成更新。