注意力不是全连接层换名字:多头自注意力的张量实验

多头自注意力常被压缩成几行公式,于是维度、缩放和掩码稍有变化就难以定位错误。本文用一份不依赖深度学习框架的 Python 实验,把投影、分头、打分、稳定 Softmax、因果掩码和合并输出逐步展开,并用可复制断言检查概率和形状。

实验问题:同一批数字究竟经过了什么

今天算法频道里"手撕 Self-Attention、Multi-head Attention"的条目靠前,说明不少读者并不满足于会调用组件。真正容易混淆的不是公式本身,而是公式与数组形状之间的对应关系。我们因此不用训练框架,也不讨论参数如何学习,只观察一次前向计算。

设序列长度为 n,模型维度为 d,头数为 h。输入矩阵 X 的形状是 n × d。三个线性投影得到 QKV,然后把最后一维切为 h 份,每份宽度 dk=d/h。第 r 个头计算:

softmax(Q_r K_r^T / sqrt(dk)) V_r

最后把所有头按特征维拼回 n × d。多头并不是把同一结果复制多遍,因为每个头拿到的是不同投影后的子空间。即使示例里为了可读性使用确定矩阵,这个结构也已经存在。

先验证缩放项

点积会随着维度增长而增大。若 QK 每个分量方差近似为 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 内容。这样从教学矩阵扩展到工程组件时,测试仍围绕算法性质,而不是围绕某个框架偶然的张量布局。

上线回归还应固定一份小输入和参数快照,在更换库版本、算子融合或精度格式后重复比较。低精度允许更宽容差,但掩码位置为零、概率行和与输出形状仍是硬约束,不能用"浮点误差"解释结构性失败。

相关推荐
Postkarte不想说话1 小时前
vLLM自定义对话模板
人工智能
具身AGI1 小时前
宇树打新:物理AI 国产的「本体」先跑通了商业化
人工智能
Json____1 小时前
AI内容创作平台项目源码
人工智能·ai·agent·内容创作·wwwoop.com
Jay-r1 小时前
DeepSeek Harness 极简上手:装好、玩熟、让它自己长新能力
人工智能·windows·ai·github·ai编程·deepseek·harness
CodeBlog-star1 小时前
LLM能力与边界:多模态、幻觉、上下文窗口及开源模型对比
人工智能·python·开源·llm
COOLMO研究AI1 小时前
Python 如何实现 AI API 的动态路由与多通道负载均衡:多账号与多供应商的高可用调度
人工智能·python·负载均衡
不懂的浪漫1 小时前
吴恩达《AI Engineering Skills Map》译读:四项核心能力与持续学习底座
人工智能·学习
冬奇Lab2 小时前
Code Agent 解剖(02):agent 是怎么一轮一轮思考和行动的?
人工智能·llm·agent
菜冻鱼2 小时前
Python-sklearn-降维
开发语言·人工智能·python·机器学习·支持向量机·sklearn