Transformer 常被缩写成一张由箭头、残差连接和矩阵乘法组成的结构图。直接记住这张图并不难,困难的是回答几个更具体的问题:一个 token 怎样变成向量,QKV 分别做什么,为什么要遮住未来,以及训练时能并行计算的模型为什么生成时仍要一步一步输出。

本文从最小神经网络单元出发,用三个 token 手算一次注意力,再把它放回完整的 Transformer 块。重点是建立一条可以检查的计算链,而不是把“注意力”描述成模型像人一样集中精神。

1. 先把 Transformer 放回神经网络

神经网络接收一组数值,通过带参数的变换产生另一组数值。最常见的线性层可以写成:

1
y = xW + b

x 是输入向量,Wb 是训练中更新的参数。若连续堆叠线性层,中间没有非线性激活,整个网络仍可合并成一次线性变换;因此实际网络会加入 ReLU、GELU 等非线性函数。

概念 在计算中是什么 容易产生的误解
参数 训练更新的矩阵或向量 参数多就必然理解得更深
激活值 当前样本经过网络时产生的中间结果 它是模型长期保存的知识
表示 用一组数值编码当前对象及上下文 每个维度都对应固定的人类概念
损失 衡量预测与训练目标差异的标量 损失较低就覆盖所有真实需求
梯度 损失对参数变化的局部敏感度 梯度下降能保证找到全局最优解

训练循环可以概括为“前向计算—计算损失—反向传播—更新参数”。Transformer 没有离开这套框架;它改变的是序列中各位置交换信息的方式。

这里的 softmax 只把一组 logits 归一化成概率。模型输出的是条件分布,不是从数据库中直接取回一句确定答案。

2. token、向量和位置缺一不可

分词器先把文本切成词、子词或字节片段,并把每个 token 映射成整数 ID。嵌入表再根据 ID 取出一个向量。相同 token 起初取得相同的基础向量,经过多层上下文计算后,各位置的表示才会因上下文不同而改变。

单看一组 token 向量,注意力运算并不知道它们的先后顺序。原始 Transformer 使用正弦与余弦位置编码;后续模型也会使用可学习位置向量、相对位置偏置或旋转位置编码。具体方案并不统一,但目的相同:让位置关系参与计算。

输入成分 回答的问题 是否随上下文改变
token ID 这是词表里的哪一个离散单元
token embedding 这个单元的基础连续表示是什么 训练后固定,推理时查表
位置信息 它处在序列哪里,与别处相隔多远 随位置改变
隐藏状态 结合当前上下文后,这个位置怎样表示 每层都会改变

位置编码不是给 token 简单贴上“第几个词”的标签。它进入后续矩阵运算,使注意力和前馈网络能够根据顺序形成不同结果。

3. 自注意力到底计算什么

设输入矩阵为 X,每一行对应一个 token。三个可训练投影把它变成:

1
2
3
4
Q = XWq
K = XWk
V = XWv
Attention(Q, K, V) = softmax(QKᵀ / √dₖ) V

可以把一行 Q 理解为当前位置发出的匹配请求,一行 K 是每个候选位置用于匹配的特征,一行 V 是匹配后真正汇总的信息。这个说法只是帮助理解矩阵职责;它不意味着网络内部存在自然语言形式的“问题”和“答案”。

除以 √dₖ 是为了控制点积的典型尺度。维度增大时,未经缩放的点积绝对值容易变大,使 softmax 很快接近极端分布,梯度也可能变得不利于训练。缩放不会保证权重均匀,它只是让数值尺度更可控。

用三个位置手算一行

假设当前查询与三个键的缩放后分数分别是 1.2、2.0、0.4。softmax 先取指数,再除以指数和:

被关注位置 分数 exp(分数) 近似值 注意力权重
1.2 3.320 0.272
2.0 7.389 0.605
0.4 1.492 0.122

三个权重之和约为 1。随后用它们对三行 V 做加权和,得到当前位置的新向量。权重最大的键贡献较大,但输出仍可能混合多个位置的信息。

亲手切换因果遮罩

下面固定了一组教学分数。选择查询位置,再切换因果遮罩,可以观察未来 token 如何在 softmax 之前被排除。数值是人为构造的示例,不是某个真实模型的内部权重。

注意力实验

同一组分数,遮罩会怎样改变归一化

选择发出查询的 token;开启因果遮罩后,它只能读取自己和左侧位置。

遮罩后的分数常用一个很大的负数表示,使其 softmax 权重趋近 0。不能先对全部位置做 softmax,再随手删掉未来权重而不重新归一化,否则剩余权重之和将小于 1。

4. 为什么需要多头注意力

单头注意力只在一组投影空间里计算关系。多头注意力把隐藏维度分给多个头,各自生成 Q、K、V,独立计算注意力后再拼接并投影:

1
2
headᵢ = Attention(XWqᵢ, XWkᵢ, XWvᵢ)
MultiHead(X) = Concat(head₁, ..., headₕ)Wo

不同头可以学习不同的匹配模式,但不能预先保证某个头必然对应“主语”或“时间”。观察到的权重模式也不能单独证明模型作出预测的因果理由。

结构 作用 没有保证的事情
多个注意力头 在多个投影子空间并行汇总信息 每个头都可被人清楚命名
输出投影 Wo 混合各头的输出 保留每个头的独立语义
dropout 训练时随机抑制部分连接 推理时仍随机丢弃同样位置

头数增多时,单头维度通常会相应变化。比较模型时应同时看隐藏维度、头数、头维度和参数共享方式,不能只用“头更多”判断能力。

5. 一个 Transformer 块不只有注意力

注意力负责位置之间交换信息,前馈网络则对每个位置分别进行相同的非线性变换。残差连接提供较短的信息与梯度路径,归一化帮助控制激活尺度。

上图采用常见的 pre-norm 顺序。2017 年原始 Transformer 论文使用的是子层之后再归一化的 post-norm 结构。阅读模型配置或源码时,应确认归一化位于子层之前还是之后,不能把所有 Transformer 块视为完全相同。

前馈网络常写成两次线性变换和一次激活:

1
FFN(x) = W₂ activation(W₁x + b₁) + b₂

同一组 FFN 参数用于每个位置,但各位置输入已经被注意力改写,因此输出仍包含上下文差异。

6. 编码器、解码器与三种常见路线

原始 Transformer 是机器翻译用的编码器—解码器结构。后来常见模型会只使用其中一侧或调整遮罩方式。

路线 注意力可见范围 常见训练目标 适合建立的直觉
编码器型 通常双向读取输入 遮盖 token、分类等 理解整段输入后产生表示
解码器型 因果遮罩,只看当前位置及之前 预测下一个 token 从左到右生成序列
编码器—解码器型 编码器双向;解码器还读取编码结果 条件序列生成 先编码输入,再生成输出

“Transformer”指一类结构原则,不等于某个固定模型,也不等于某个软件库。Hugging Face 的 transformers 是实现和使用许多模型的工具库;名称相同,概念层级不同。

7. 训练能并行,生成为什么仍是串行的

训练解码器模型时,已知完整训练序列,可以把输入右移一位,并用因果遮罩同时计算各位置的下一个 token 损失。例如输入“春 天 来”,目标可以是“天 来 了”。每个位置只读取左侧,但这些位置的矩阵运算能够一起执行。

推理时,下一个 token 尚未存在。模型必须先生成第一个新 token,把它追加到上下文,再计算下一个。因此单条自回归序列在 token 维度上具有依赖链。

阶段 已知内容 可并行的部分 主要限制
训练 完整输入与目标 token 批次、序列位置和矩阵运算 显存、数据吞吐、反向传播
首次提示处理 整段提示 提示内多个位置 上下文长度与注意力计算
自回归生成 只有已生成前缀 层、头、批次中的矩阵计算 新 token 之间必须按顺序产生

KV Cache 会保存先前 token 各层的键和值,新一步只需计算新 token 的 Q、K、V,再让新查询读取缓存。它避免重复计算旧位置的键和值,但缓存占用随序列增长,新查询仍需读取越来越多的历史位置。

8. 基础注意力的时间与空间代价

长度为 n 的序列会形成 n × n 注意力分数矩阵。基础全注意力的这部分计算和中间存储随 增长;当上下文翻倍时,矩阵元素数量变为四倍。

FlashAttention 等实现通过分块和改善内存访问减少中间数据搬运,并能得到精确注意力结果,但它没有把数学上所有查询—键配对自动变成线性数量。滑动窗口、稀疏注意力等方法则会改变可见连接模式,需要另行分析表达能力和边界行为。

优化方法 主要减少什么 仍要确认什么
KV Cache 生成时对历史 K、V 的重复计算 缓存显存、批处理与长序列读取成本
融合注意力内核 中间张量和显存读写 硬件、数据类型与后端支持
局部或稀疏注意力 实际参与匹配的位置数量 被省略的长距离关系是否重要
量化 权重或缓存的存储与带宽 精度变化和算子支持

工程文档中的“更快”必须连同硬件、形状、数据类型、批大小和精度一起阅读。不同注意力后端可能在浮点舍入上产生细小差异。

9. 用纯 Python 跑通缩放点积注意力

下面只使用标准库,实现矩阵乘法、稳定 softmax 和可选因果遮罩。它适合核对公式,不用于替代 PyTorch 等框架的高性能算子。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
from math import exp, sqrt
import json


def dot(left, right):
return sum(a * b for a, b in zip(left, right))


def scaled_dot_product_attention(query, key, value, causal=False):
if not query or not key or not value:
raise ValueError("Q、K、V 不能为空")
if len(key) != len(value):
raise ValueError("K 和 V 的位置数必须相同")
key_dim = len(key[0])
value_dim = len(value[0])
if key_dim == 0 or value_dim == 0:
raise ValueError("向量维度必须为正")
if any(len(row) != key_dim for row in key):
raise ValueError("K 的每一行维度必须相同")
if any(len(row) != key_dim for row in query):
raise ValueError("Q 与 K 的向量维度必须相同")
if any(len(row) != value_dim for row in value):
raise ValueError("V 的每一行维度必须相同")
if causal and len(query) != len(key):
raise ValueError("这个简化实现只对等长自注意力使用因果遮罩")

scale = sqrt(key_dim)
all_weights = []
outputs = []

for query_index, query_row in enumerate(query):
scores = []
for key_index, key_row in enumerate(key):
if causal and key_index > query_index:
scores.append(float("-inf"))
else:
scores.append(dot(query_row, key_row) / scale)

largest = max(score for score in scores if score != float("-inf"))
numerators = [0.0 if score == float("-inf") else exp(score - largest)
for score in scores]
denominator = sum(numerators)
weights = [number / denominator for number in numerators]
output = [sum(weights[j] * value[j][d] for j in range(len(value)))
for d in range(value_dim)]
all_weights.append(weights)
outputs.append(output)

return all_weights, outputs


if __name__ == "__main__":
token_vectors = [[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]]
weights, outputs = scaled_dot_product_attention(
token_vectors, token_vectors, token_vectors, causal=True
)
print(json.dumps({"weights": weights, "outputs": outputs}, ensure_ascii=False))

运行后可以检查三个不变量:每行权重之和约为 1,被因果遮罩的位置权重为 0,每个输出维度都是 V 对应维度的加权和。真实模型还会处理批次、多头、padding mask、dropout 和不同数据类型。

10. 阅读模型代码时按形状追踪

第一次读实现时,先不要陷入类名。为每个张量标出形状,很多错误会立即显现。下面省略批次前的其他维度:

张量 典型形状 关键检查
隐藏状态 X [batch, length, d_model] token 与特征维是否放反
分头后的 Q [batch, heads, q_len, d_head] d_model 如何拆到各头
分头后的 K [batch, kv_heads, kv_len, d_head] 是否使用多查询或分组查询注意力
注意力分数 [batch, heads, q_len, kv_len] mask 能否正确广播
注意力输出 [batch, heads, q_len, d_value] 拼接前是否恢复正确顺序

不同 API 对布尔遮罩的语义可能相反。例如 PyTorch 当前文档说明,scaled_dot_product_attention 的布尔遮罩中 True 表示参与注意力,而 MultiheadAttentionkey_padding_maskTrue 表示需要屏蔽。移植代码时必须查当前接口文档,不能只凭变量名猜测。

11. 常见误区与排查清单

  1. 把 token 当作完整单词。 子词和字节级分词都可能把一个词拆成多个单元。
  2. 忘记位置。 没有位置信息时,自注意力本身不能区分同一组 token 的排列。
  3. Q、K、V 当成三份原文。 它们是输入经过不同可训练矩阵后的表示。
  4. 在 softmax 后错误处理遮罩。 删除权重后必须保证剩余项重新归一化。
  5. 把注意力权重直接当解释。 权重展示了一次信息混合比例,不足以单独给出因果解释。
  6. 认为训练和生成完全同速。 训练已知目标序列,生成必须等待前一个新 token。
  7. 认为 KV Cache 消除了长上下文成本。 它减少重复计算,也增加持续增长的缓存。
  8. 忽略归一化顺序。 pre-norm 与 post-norm 的计算图和训练性质不同。
  9. 只看模型名称判断结构。 应检查是否是编码器、解码器或编码器—解码器,以及具体 mask。
  10. 手写生产注意力内核。 教学实现用于理解;部署时应优先使用框架经过测试的算子。

12. 原始资料与下一步

资料核对日期:2026-09-17。

  1. Attention Is All You Need:原始 Transformer 论文,适合核对缩放点积注意力、多头结构、位置编码和编码器—解码器设计。
  2. PyTorch scaled_dot_product_attention 文档:核对当前函数签名、张量形状、遮罩语义与可用后端。
  3. Hugging Face Transformers:Attention backends:理解模型接口怎样选择 SDPA、FlashAttention 等实现,以及自定义后端时为何必须同时处理遮罩。

理解这篇文章后,可以沿两条路线继续:想理解模型怎样通过反馈学习行为,可读 强化学习入门;想理解模型拿到外部证据后怎样被系统化评测,可读 RAG 知识库评测。最后再回到 大模型推理时计算,区分模型内部的一次前向计算与系统在外部增加采样、验证和修订预算。