Agent 的一切都跑在注意力机制之上。这一章只做一件事:让你能在白板上画出从 token 到注意力矩阵的形状变化,并说清复杂度里那个 n² 究竟从哪来——因为后面每一章的成本、延迟、上下文策略,都是这个 n² 的推论。
编辑部注 · 本章不推导公式,只走一遍数据形状。读完后请合上文章,自己在纸上画一遍 Q/K/V 的维度。
如果要给「Harness 工程师为什么要懂模型原理」找一个最实际的理由,那就是:你所有的取舍最后都会被这三个物理事实定死——注意力是平方复杂度的、KV Cache 是按前缀命中的、输出是采样出来的。这一章讲第一个。
一、为什么一个做工程的人要关心注意力
面试里被问到「Transformer 原理」,很多人以为考官在考古。其实不是。考官真正想问的是:
- 你知不知道上下文变长,成本是怎么长的?
- 你知不知道「把整个代码库塞进上下文」这条路会在哪一步崩掉?
- 你知不知道长任务的瓶颈到底在算力、显存,还是在带宽?
这三个问题都从自注意力出发。凡是能用 n² 解释清楚的现象,都不该用「模型能力不够」来解释。
二、一次前向到底做了什么
先不碰公式,只跟踪张量的形状。设一批输入是 n 个 token,每个 token 进模型后被映射成一个 d 维向量(例如 d=4096)。
[n, d] 的输入先压成 [n, n] 的权重矩阵,再用它加权聚合 V。所有关于「上下文很贵」的直觉,都来自中间那个 [n, n]:它随 n 平方增长,而输入本身只随 n 线性增长。三步走,记住这三步就够用了:
- 投影:同一份输入 X,用三组不同的权重矩阵投影出 Q(我想找什么)、K(我能被什么找到)、V(我实际携带的内容)。
- 打分:用
Q · Kᵀ算出任意两个位置之间的相关度,得到一个[n, n]的矩阵,除以√d做缩放,再 softmax 成权重。 - 聚合:用这组权重对 V 做加权求和,输出形状回到
[n, d]。
三、n² 究竟出在哪一步
把每一步的复杂度摊开看,问题一目了然。
| 阶段 | 做的事 | 张量形状 | 计算量级 | 随 n 的增长 |
|---|---|---|---|---|
| 投影 | X·W_q/k/v | [n, d] → [n, d] | O(n · d²) | 线性 |
| 打分 | Q·Kᵀ | [n, d] × [d, n] → [n, n] | O(n² · d) | 平方 |
| softmax | 按行归一化 | [n, n] | O(n²) | 平方 |
| 聚合 | A·V | [n, n] × [n, d] → [n, d] | O(n² · d) | 平方 |
| 前馈网络 | 逐位置 MLP | [n, d] → [n, d] | O(n · d²) | 线性 |
结论:平方项只来自「任意两个位置都要互相看一眼」这件事,也就是那两步 QKᵀ 和 AV。剩下的一切(嵌入、投影、FFN)都是逐位置的,随 n 线性增长。
设 d = 4096,n = 8,000(大约 3 万汉字或 6,000 行代码):
- 线性项的量级约
n·d² ≈ 1.34×10¹¹; - 平方项的量级约
n²·d ≈ 2.6×10¹¹。
此时两者还在同一数量级——这就是「中等长度上下文时,注意力还不是绝对瓶颈」的原因。但只要 n 再翻一倍,平方项就翻四倍,线性项只翻两倍。拐点大约出现在 n 与 d 可比的时候;一旦 n 远超 d,平方项就彻底主导成本。
四、这对做 Agent 意味着什么
把上面的结论翻译成工程语言,会得到三条非常硬的推论:
(1) 成本不是线性的,所以「多塞点上下文」是一种高息负债。 往上下文里加一倍内容,注意力部分的代价接近四倍。这解释了为什么长上下文模型的 API 定价通常按阶梯上涨,而不是线性。
(2) 裁剪与压缩不是优化,而是必备能力。 如果你的 Harness 只会「一直追加消息」,那你实际上是在让成本做平方增长。上下文压缩、状态外置、历史摘要这三件事,是 Harness 的基本功而不是加分项。
(3) 首 token 延迟与后续 token 延迟是两个不同的成本。 处理输入(prefill)要跑完整个 [n, n];生成每个 token(decode)只走一步,但需要读缓存。这两个阶段的瓶颈不同,优化手段也不同——这是下一章 KV Cache 的主题。
把 n² 记住,你就不需要背任何「长上下文为什么要小心」的结论了——那些结论全都是它的推论。本刊编辑部
五、动手:三十行手写一遍注意力
不写一遍就容易把形状记混。下面这段 NumPy 代码把图 1 完整实现了一遍,重点看注释里标注的形状。
import numpy as np
def softmax(x, axis=-1):
m = x.max(axis=axis, keepdims=True) # 数值稳定:先减去最大值
e = np.exp(x - m)
return e / e.sum(axis=axis, keepdims=True)
def attention(X, Wq, Wk, Wv):
"""X: [n, d] 三个权重: [d, d]"""
Q = X @ Wq # [n, d] 我想找什么
K = X @ Wk # [n, d] 我能被什么找到
V = X @ Wv # [n, d] 我携带什么
d = Q.shape[-1]
S = Q @ K.T # [n, n] ← 平方复杂度的来源
S = S / np.sqrt(d) # 缩放,防止 softmax 进入饱和区
A = softmax(S) # [n, n] 每行和为 1
return A @ V # [n, d] 回到原形状
n, d = 5, 8
rng = np.random.default_rng(0)
X = rng.normal(size=(n, d))
Wq = rng.normal(size=(d, d)) / np.sqrt(d)
Wk = rng.normal(size=(d, d)) / np.sqrt(d)
Wv = rng.normal(size=(d, d)) / np.sqrt(d)
out = attention(X, Wq, Wk, Wv)
print(out.shape) # (5, 8) —— 形状与输入一致,这是所有注意力变体的共同特征
# 顺手验证因果掩码(decoder 只能看左边)
mask = np.triu(np.ones((n, n)), k=1).astype(bool)
print(mask.astype(int))- 把
n从 5 改到 2000,测量S = Q @ K.T与X @ Wq的耗时,亲手看到平方项吃掉时间; - 把
mask加进S(被 mask 的位置填-inf),再做 softmax,确认第一行只依赖第一个 token; - 回答自己一个问题:如果我想让「第 500 个 token 只关注最近 100 个 token」(滑窗注意力),代码要改哪一行?
六、自测
QKᵀ 产生 [n, n]),才出现平方项。记住这个区分,就能判断任何「优化长上下文」的技术到底在优化什么。七、小结与术语表
| 术语 | 一句话解释 | 工程含义 |
|---|---|---|
| Self-Attention | 序列中每个位置对其他所有位置算权重并加权聚合 | 产生 [n, n] 中间矩阵,是平方复杂度的唯一来源 |
| Q / K / V | 同一输入经三组权重投影出的查询、键、值 | 形状都是 [n, d];K、V 正是后续会被缓存起来的对象 |
| 缩放因子 √d | 点积结果除以维度的平方根 | 防止 softmax 进入饱和区导致梯度消失与权重退化成 one-hot |
| 因果掩码 | 让位置 i 只能看到 i 及之前的位置 | 这是 decoder-only 模型可做 KV Cache 的前提 |
| Prefill / Decode | 处理输入的阶段 / 逐个生成 token 的阶段 | 瓶颈不同:前者算力受限,后者显存带宽受限 |
下一章我们把这个 n² 换个角度再看一次:既然前缀会被反复重算,那能不能缓存?能,但缓存命中与否取决于一个你平时根本不会注意到的东西——前缀稳定性。
◇ 面试官会怎么问
共 3 条 · 其中 1 条高频 · 先自己答一遍,再展开对照
为什么注意力复杂度是 O(n²)?这对做 Agent 意味着什么? 高频
seq × seq 的注意力矩阵,所以算力随序列长度呈平方增长。对 Agent 的直接含义有三层:第一,上下文越长,成本不是线性上涨而是平方级上涨,这是"长上下文很贵"的物理原因;第二,往上下文里多塞一倍内容,代价远不止翻一倍;第三,裁剪、压缩、状态外置不是优化技巧,而是必须做的基本功。
追问链
- 那 KV Cache 的显存占用怎么估算?它和 n² 是同一个量级吗?
- 如果只关心首 token 延迟,哪个成本占主导?
多头注意力比单头好在哪?
工程意义在于:模型能同时维持多种关系,这是长链路任务里"既记住目标又盯住细节"的基础。注意头的数量与 head_dim 的乘积大致等于隐藏维度,所以加头并不等比例增加算力,但头太多也会稀释每个头的表达空间。
位置编码解决什么问题?没有会怎样?
对工程的启示:相对位置方案让模型对超出训练长度的位置有一定外推能力,但外推能力有限。"把窗口调大"不等于"模型真的能用好长窗口",这就是长上下文需要专门训练和评测的原因。
结论先行(一句话给出取舍)→ 机制(为什么是这样,涉及哪条链路)→ 代价或边界(这么做放弃了什么)→ 你的实践或数字(真实场景里怎么落地)。 只讲机制不讲代价,是背题;只讲取舍没有数字,是空谈。