多头自注意力常被压缩成几行公式,于是维度、缩放和掩码稍有变化就难以定位错误。本文用一份不依赖深度学习框架的 Python 实验,把投影、分头、打分、稳定 Softmax、因果掩码和合并输出逐步展开,并用可复制断言检查概率和形状。
实验问题:同一批数字究竟经过了什么
今天算法频道里"手撕 Self-Attention、Multi-head Attention"的条目靠前,说明不少读者并不满足于会调用组件。真正容易混淆的不是公式本身,而是公式与数组形状之间的对应关系。我们因此不用训练框架,也不讨论参数如何学习,只观察一次前向计算。
设序列长度为 n,模型维度为 d,头数为 h。输入矩阵 X 的形状是 n × d。三个线性投影得到 Q、K、V,然后把最后一维切为 h 份,每份宽度 dk=d/h。第 r 个头计算:
softmax(Q_r K_r^T / sqrt(dk)) V_r
最后把所有头按特征维拼回 n × d。多头并不是把同一结果复制多遍,因为每个头拿到的是不同投影后的子空间。即使示例里为了可读性使用确定矩阵,这个结构也已经存在。
先验证缩放项
点积会随着维度增长而增大。若 Q、K 每个分量方差近似为 1,独立点积的方差约为 dk。除以 sqrt(dk) 后,打分尺度回到较稳定的范围。缺少缩放时,Softmax 很容易接近 one-hot,微小输入扰动会产生过于剧烈的概率变化。这里的缩放不是为了让结果"更平均",而是为了控制数值尺度。
Softmax 本身还要先减去每行最大值。数学上,所有指数同时乘同一常数不会改变归一化结果;工程上,这一步避免 exp(1000) 溢出。若某位置被因果掩码禁止访问,我们把它视为负无穷,使其指数为零。
完整 Python 实验
下面只用标准库实现矩阵操作。权重采用循环移位矩阵,使每个头能看到不同组合,又方便复算。程序同时打印第一头的注意力矩阵,并验证每行概率和为 1、未来位置概率为 0、最终形状不变。
python
import math
def matmul(a, b):
rows, mid, cols = len(a), len(b), len(b[0])
assert all(len(row) == mid for row in a)
assert all(len(row) == cols for row in b)
return [[sum(a[i][k] * b[k][j] for k in range(mid))
for j in range(cols)] for i in range(rows)]
def transpose(a):
return [list(col) for col in zip(*a)]
def softmax(row):
finite = [x for x in row if x != float("-inf")]
peak = max(finite)
exps = [0.0 if x == float("-inf") else math.exp(x - peak)
for x in row]
total = sum(exps)
return [x / total for x in exps]
def project(x, shift):
d = len(x[0])
w = [[0.0] * d for _ in range(d)]
for i in range(d):
w[i][(i + shift) % d] = 1.0
return matmul(x, w)
def split_heads(x, heads):
width = len(x[0]) // heads
return [[[row[h * width + j] for j in range(width)]
for row in x] for h in range(heads)]
def attention(x, heads=2, causal=True):
n, d = len(x), len(x[0])
assert d % heads == 0
q_heads = split_heads(project(x, 0), heads)
k_heads = split_heads(project(x, 1), heads)
v_heads = split_heads(project(x, 2), heads)
width = d // heads
head_outputs, all_weights = [], []
for q, k, v in zip(q_heads, k_heads, v_heads):
raw = matmul(q, transpose(k))
weights = []
for i, row in enumerate(raw):
scores = []
for j, value in enumerate(row):
blocked = causal and j > i
scores.append(float("-inf") if blocked
else value / math.sqrt(width))
weights.append(softmax(scores))
all_weights.append(weights)
head_outputs.append(matmul(weights, v))
output = []
for token in range(n):
merged = []
for h in range(heads):
merged.extend(head_outputs[h][token])
output.append(merged)
return output, all_weights
if __name__ == "__main__":
x = [[1.0, 0.0, 1.0, 0.0],
[0.0, 2.0, 0.0, 1.0],
[1.0, 1.0, 0.0, 0.0]]
out, weights = attention(x, heads=2, causal=True)
assert len(out) == 3 and all(len(row) == 4 for row in out)
for matrix in weights:
for i, row in enumerate(matrix):
assert abs(sum(row) - 1.0) < 1e-9
assert all(abs(row[j]) < 1e-12 for j in range(i + 1, 3))
for row in weights[0]:
print(" ".join(f"{v:.4f}" for v in row))
测试输入就是代码中的三枚 token。运行后第一行只能看自己,所以为 1.0000 0.0000 0.0000;第二行第三列仍为零;第三行可以访问全部位置。断言比固定整张浮点结果更稳健,因为它验证的是算法必须保持的不变量。
一次改一个变量
把 causal 改为 False,第一行不再只看自己,这对应编码器的双向注意力。把 heads 改成 4,每头宽度变成 1,仍能运行;改成 3 则会触发整除断言。真实模型还会在合并后增加输出投影 W_O,但它不改变分头与聚合的核心逻辑。
实验也揭示一个常见误解:掩码不是把输出位置删除,而是改变权重归一化的候选集合。被屏蔽位置必须在 Softmax 前处理。若先求概率再把未来位置清零,剩余概率和会小于 1,输出尺度随可见位置数量漂移。
复杂度与内存账单
投影若使用稠密权重,时间为 O(n d²);每个头的打分与加权总计为 O(n² d)。注意力矩阵需要 O(h n²) 空间,通常是长序列的主要瓶颈。本示例的移位投影仍用通用矩阵乘法,便于保持结构清晰;生产实现会用优化内核,并可能分块计算以避免完整保存打分矩阵。
边界与失败记录
空序列应在调用前拒绝,否则无法推断维度。d 必须能被头数整除;每行输入宽度必须一致;因果掩码至少要保留当前位置,否则一整行全为负无穷,Softmax 没有分母。大数输入必须使用减最大值的稳定实现。最后,浮点测试不要直接比较字符串或要求完全相等,应使用容差并检查行和、非负性和掩码位置。
常见错误还包括把 K 忘记转置、按 token 维而不是特征维切头、合并时交错顺序错误,以及把缩放因子写成 sqrt(d)。这些错误有时不会报维度异常,却会改变模型含义,所以形状断言与性质断言都不可少。
实验结论
多头自注意力可以拆成六个可单测步骤:投影、切头、点积、缩放与掩码、稳定归一化、合并。公式简短不代表实现可以省略不变量。先用小矩阵看清每个位置能访问谁,再进入框架和 GPU 内核,定位问题会快得多。
从小实验迁移到批量实现
真实组件通常还多一个批次维,形状从 n×d 变为 batch×n×d。常见库会把三个投影合并成一次大矩阵乘法,再重排为 batch×heads×n×width。这只是减少内核调用,不改变本文六步逻辑。迁移时最好在重排前后写出形状表,并用一个批次、一个头、一个 token 逐级退化测试;当维度都为一时,许多错误广播反而会被隐藏,因此还要补一个各维长度都不同的小样本。
批次掩码也有两类。因果掩码由位置关系决定,形状通常可广播到所有批次和头;填充掩码由每条样本的有效长度决定,不同批次并不相同。两者合并后,任何查询行至少要保留一个合法键。若填充位置本身也发起查询,调用方还要在输出阶段清零或忽略它,否则即使键侧被遮住,查询侧仍会产生一个归一化向量。
用不变量设计更多测试
除了检查概率行和,还可以构造全零输入。若投影没有偏置,QK^T 全为零;非因果注意力应在所有合法位置均匀分配。再构造两个完全相同的 token,在没有位置编码时交换它们,输出也应按同样方式交换。这种置换等变性是自注意力的结构属性,若测试失败,往往是切头、拼接或掩码索引混入了绝对位置。
梯度不在本文代码范围内,但前向性质仍能帮助框架实现。可以把标准库版本当作小尺寸参考,与框架张量逐元素比较。测试权重要显式固定,关闭 dropout,并确保浮点类型相同。差异若随序列长度放大,先检查 Softmax 稳定化与累加顺序;若只在多个头时出现,优先检查重排和拼接轴。
长序列不是简单换个循环
当 n 翻倍,注意力打分元素数量扩大四倍。分块或流式注意力通过维护每行的局部最大值、指数和与加权和值,逐块合并稳定 Softmax,避免完整落地 n×n 矩阵。合并时不能直接把各块归一化后的输出相加,因为每块分母不同;必须用全局最大值重新缩放局部统计量。这也是为什么高性能实现虽改变内存路径,却必须保持数学不变量。
稀疏注意力、滑动窗口注意力和低秩近似进一步改变候选集合或表示方式,复杂度可能降低,但模型语义也随之变化。评估时要同时记录吞吐、峰值内存与任务质量,不能只用一段随机矩阵的运行时间替代实际效果。标准全注意力的小矩阵实现仍值得保留,它是验证近似版本与优化内核的基准裁判。
发布前核对表
首先核对输入维度、头数与每头宽度;其次确认掩码在 Softmax 前生效,且不会产生全屏蔽行;然后检查数值稳定化、概率非负和行和;最后验证输出投影、残差连接与归一化层属于组件的哪一层,避免重复执行。日志只记录形状、范围和异常计数,不打印真实 token 内容。这样从教学矩阵扩展到工程组件时,测试仍围绕算法性质,而不是围绕某个框架偶然的张量布局。
上线回归还应固定一份小输入和参数快照,在更换库版本、算子融合或精度格式后重复比较。低精度允许更宽容差,但掩码位置为零、概率行和与输出形状仍是硬约束,不能用"浮点误差"解释结构性失败。