Agent 工程学习站 做棵大树 出品 · beatree.cn 免费 · 无需登录 · 进度存本机
I 版 · 大模型原理 第 01 章 Attention Is All You Need
I · 大模型原理 01 / 22 入门

Transformer 与自注意力

Attention Is All You Need
预计阅读 22 分钟
难度 入门
关键词 Self-Attention · QKV
本机状态 未读

Agent 的一切都跑在注意力机制之上。这一章只做一件事:让你能在白板上画出从 token 到注意力矩阵的形状变化,并说清复杂度里那个 n² 究竟从哪来——因为后面每一章的成本、延迟、上下文策略,都是这个 n² 的推论。

编辑部注 · 本章不推导公式,只走一遍数据形状。读完后请合上文章,自己在纸上画一遍 Q/K/V 的维度。

如果要给「Harness 工程师为什么要懂模型原理」找一个最实际的理由,那就是:你所有的取舍最后都会被这三个物理事实定死——注意力是平方复杂度的、KV Cache 是按前缀命中的、输出是采样出来的。这一章讲第一个。

一、为什么一个做工程的人要关心注意力

面试里被问到「Transformer 原理」,很多人以为考官在考古。其实不是。考官真正想问的是:

  • 你知不知道上下文变长,成本是怎么长的?
  • 你知不知道「把整个代码库塞进上下文」这条路会在哪一步崩掉?
  • 你知不知道长任务的瓶颈到底在算力、显存,还是在带宽?

这三个问题都从自注意力出发。凡是能用 n² 解释清楚的现象,都不该用「模型能力不够」来解释。

二、一次前向到底做了什么

先不碰公式,只跟踪张量的形状。设一批输入是 n 个 token,每个 token 进模型后被映射成一个 d 维向量(例如 d=4096)。

自注意力:只看形状 n = 序列长度(token 数) d = 隐藏维度(如 4096) h = 头数 输入 token 序列 形状 [n] Embedding + 位置编码 形状 [n, d] Q = X · W_q [n, d] K = X · W_k [n, d] V = X · W_v [n, d] A = softmax(QKᵀ / √d) 形状 [n, n] ← 平方复杂度的源头 注意力矩阵(每格是一次 token 间权重) n 列 × n 行 = n² 个数 输出 = A · V 形状 [n, d](回到原尺寸)
图 1 注意力的核心只有一步:把 [n, d] 的输入先压成 [n, n] 的权重矩阵,再用它加权聚合 V。所有关于「上下文很贵」的直觉,都来自中间那个 [n, n]它随 n 平方增长,而输入本身只随 n 线性增长。

三步走,记住这三步就够用了:

  1. 投影:同一份输入 X,用三组不同的权重矩阵投影出 Q(我想找什么)、K(我能被什么找到)、V(我实际携带的内容)。
  2. 打分:用 Q · Kᵀ 算出任意两个位置之间的相关度,得到一个 [n, n] 的矩阵,除以 √d 做缩放,再 softmax 成权重。
  3. 聚合:用这组权重对 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 完整实现了一遍,重点看注释里标注的形状。

scaled_dot_product_attention.pypython
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))
实操任务(建议 20 分钟)
  • n 从 5 改到 2000,测量 S = Q @ K.TX @ Wq 的耗时,亲手看到平方项吃掉时间;
  • mask 加进 S(被 mask 的位置填 -inf),再做 softmax,确认第一行只依赖第一个 token;
  • 回答自己一个问题:如果我想让「第 500 个 token 只关注最近 100 个 token」(滑窗注意力),代码要改哪一行?

六、自测

本章自测答错的题建议加入复习队列
Q1注意力里随序列长度平方增长的是哪一部分?
B。投影与 FFN 都是逐位置计算的,复杂度 O(n·d²),随 n 线性。只有当两个位置必须互相比较时(QKᵀ 产生 [n, n]),才出现平方项。记住这个区分,就能判断任何「优化长上下文」的技术到底在优化什么。
Q2如果把上下文长度从 8k 提到 32k,注意力部分的算力大约变成原来的几倍?
C。n 变成 4 倍,n² 变成 16 倍。这正是「长上下文很贵」的物理来源,也是阶梯定价的合理性所在。注意这是注意力部分的量级,端到端延迟还受 KV Cache 读写带宽影响,实际倍率不一定正好是 16 倍,但量级关系成立。
Q3位置编码如果不加,模型会出什么问题?
B。注意力本身对输入顺序不敏感——打乱输入,输出只会跟着打乱,模型没有任何机制知道谁在前谁在后。位置编码(现在主流是 RoPE 这类相对位置方案)把顺序信息注入表示。顺带的推论是:相对位置方案让模型对超出训练长度的位置有一定外推能力,但「把窗口调大」不等于「模型真的能用好长窗口」——这是长上下文必须单独评测的原因。

七、小结与术语表

术语一句话解释工程含义
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 的直接含义有三层:第一,上下文越长,成本不是线性上涨而是平方级上涨,这是"长上下文很贵"的物理原因;第二,往上下文里多塞一倍内容,代价远不止翻一倍;第三,裁剪、压缩、状态外置不是优化技巧,而是必须做的基本功。

追问链

  1. 那 KV Cache 的显存占用怎么估算?它和 n² 是同一个量级吗?
  2. 如果只关心首 token 延迟,哪个成本占主导?
加分点:能把「注意力矩阵的平方级算力」和「KV Cache 的线性级显存」分开讲,就说明不是背结论。
多头注意力比单头好在哪?
单头只有一组 Q/K/V 投影,只能学到一种关注方式。多头把表示空间切成若干子空间并行做注意力,不同的头可以各自关注不同类型的关联——位置邻近、指代关系、句式结构等,最后拼接再投影。
工程意义在于:模型能同时维持多种关系,这是长链路任务里"既记住目标又盯住细节"的基础。注意头的数量与 head_dim 的乘积大致等于隐藏维度,所以加头并不等比例增加算力,但头太多也会稀释每个头的表达空间。
加分点:主动补一句"不是越多越好",显示你有工程判断而不是背书。
位置编码解决什么问题?没有会怎样?
注意力本身是置换等变的——把输入顺序打乱,输出只会跟着打乱,模型无法区分"猫追狗"和"狗追猫"。位置编码把顺序信息注入表示。演化路径大致是可学习绝对位置嵌入 → 正弦编码 → 现在主流的 RoPE 这类相对位置方案。
对工程的启示:相对位置方案让模型对超出训练长度的位置有一定外推能力,但外推能力有限。"把窗口调大"不等于"模型真的能用好长窗口",这就是长上下文需要专门训练和评测的原因。
加分点:把位置编码与"有效上下文长度"联系起来,是一个很自然的加分项。
答这类题的通用结构

结论先行(一句话给出取舍)→ 机制(为什么是这样,涉及哪条链路)→ 代价或边界(这么做放弃了什么)→ 你的实践或数字(真实场景里怎么落地)。 只讲机制不讲代价,是背题;只讲取舍没有数字,是空谈。

本章收尾 · Wrap up
掌握度自评
点一下给自己打分;低于 3 分建议加入复习队列
个人笔记 · Notes