摘要:多头自注意力常被压缩成几行公式,于是维度、缩放和掩码稍有变化就难以定位错误。本文用一份不依赖深度学习框架的 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、最终形状不变。
importmathdefmatmul(a,b):rows,mid,cols=len(a),len(b),len(b[0])assertall(len(row)==midforrowina)assertall(len(row)==colsforrowinb)return[[sum(a[i][k]*b[k][j]forkinrange(mid))forjinrange(cols)]foriinrange(rows)]deftranspose(a):return[list(col)forcolinzip(*a)]defsoftmax(row):finite=[xforxinrowifx!=float("-inf")]peak=max(finite)exps=[0.0ifx==float("-inf")elsemath.exp(x-peak)forxinrow]total=sum(exps)return[x/totalforxinexps]defproject(x,shift):d=len(x[0])w=[[0.0]*dfor_inrange(d)]foriinrange(d):w[i][(i+shift)%d]=1.0returnmatmul(x,w)defsplit_heads(x,heads):width=len(x[0])//headsreturn[[[row[h*width+j]forjinrange(width)]forrowinx]forhinrange(heads)]defattention(x,heads=2,causal=True):n,d=len(x),len(x[0])assertd%heads==0q_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=[],[]forq,k,vinzip(q_heads,k_heads,v_heads):raw=matmul(q,transpose(k))weights=[]fori,rowinenumerate(raw):scores=[]forj,valueinenumerate(row):blocked=causalandj>i scores.append(float("-inf")ifblockedelsevalue/math.sqrt(width))weights.append(softmax(scores))all_weights.append(weights)head_outputs.append(matmul(weights,v))output=[]fortokeninrange(n):merged=[]forhinrange(heads):merged.extend(head_outputs[h][token])output.append(merged)returnoutput,all_weightsif__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)assertlen(out)==3andall(len(row)==4forrowinout)formatrixinweights:fori,rowinenumerate(matrix):assertabs(sum(row)-1.0)<1e-9assertall(abs(row[j])<1e-12forjinrange(i+1,3))forrowinweights[0]:print(" ".join(f"{v:.4f}"forvinrow))测试输入就是代码中的三枚 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 内容。这样从教学矩阵扩展到工程组件时,测试仍围绕算法性质,而不是围绕某个框架偶然的张量布局。
上线回归还应固定一份小输入和参数快照,在更换库版本、算子融合或精度格式后重复比较。低精度允许更宽容差,但掩码位置为零、概率行和与输出形状仍是硬约束,不能用“浮点误差”解释结构性失败。