Transformer 前向传播
21 篇 primer 把 self-attention、multi-head attention、位置编码、block、以及三种架构形态,当作各自独立的机制造了出来。这一篇把它们合到一起 —— 一句话,6 个阶段,从头走到尾 —— 然后接着往前走,走到它们谁都没管过的地方:采样出一个 token、把它喂回去、让这件事划算的那个cache,以及整台机器实际跑起来到底要花多少钱。
整个 stack 的形状
21 篇 primer 造好了零件。这里它们合成一台机器 —— 6 个阶段,一句话,从头走到尾。
阶段永远是这几个:分词、embedding 加位置、跑 N 份同一个 block、归一化、反嵌入、读出 logits。attention 就住在 block 里面;它前后的一切都是记账 —— 而形状恰恰就藏在这些记账里,所以我们从这里开始。
一个 token 还不是意义 —— 它是一个 id,而这个 id就是一张表里的行号,每个不同的词占一行。拖着走过这句话,看这次查表落在哪里:
注意位置 0 和都是「the」,两者落在完全相同的一行 —— 同样的 id、同样的向量,完全不知道自己在句子里的第几个位置。这是故意的:顺序此刻还没进来。它下一步才会出现,也是位置编码这一篇的全部主题。
拉远一点,6 个阶段就是一张图:当前阶段点亮,它之前的阶段已经跑完。走一遍整个前向传播:
注意 block 那一段被特意画得更高 —— 它不是一个阶段,而是 N 份一模一样的 block 叠在一起,N 正是我们下一步要转的旋钮。self-attention、multi-head attention 和位置编码造的东西,全都装在那一段更宽的带子里。
转动这个旋钮,看参数量怎么变 —— 宽度固定在 GPT-2 small 自己的 768,只让深度动:
GPT-2 small 自己的 12 层落在 1.24 亿参数 —— 直接从曲线上读,也正是论文报告的数字。全靠深度就到了这里:词表没变,宽度没变,只是十二份 700 万参数的 block,架在一张共享的 embedding 表上。推到,参数量大约翻三倍 —— 宽度没变,曲线还是一条直线。
这 N 个 block 里的每一个,拿到手的都是残差流,能做的只有两件事:加进去,或者换掉它。切换看看:
看「换掉」做了什么:stream 自己的宽度从没变过 —— 还是 4 个格子那么宽 —— 但更早的 block 写下的东西全没了。这就是那条值得写成断言的不变式:每个层的边界上,stream 都保持自己的形状,而每个子层永远只是往里面加。§02 会把 block 打开,看看到底加了什么。
一个 block,看两遍
这 N 份里的每一份都跑同样的 4 步。走一遍就是这一节的全部 —— 机制本身在别的 primer 里已经讲过了。
一个 block 读stream、把它的一份拷贝归一化、在这份拷贝上跑attention、再把结果加回去。然后对FFN 再做一遍。两个子层,都裹在 §01 刚证明过的「加、不换」这个形状里。
当前这一步随着你走而点亮;正中间那条直线就是stream 本身,它从不拐弯 —— 每条支路都是从它岔出去,再并回来:
注意两条支路结尾都一样:attention 加一次,FFN 再加一次。这正是 §01 那条不变式,一个 block 里出现了两次 —— 这里没有任何东西会替换stream。
这两个子层干的不是一回事。挑一个 token,看看每个子层能读到什么:attention能读到其它任何位置;FFN只能读它自己坐的那个位置:
只要离开 token 0,attention 的连线就会散向全部 4 个格子,而 FFN 的连线永远不出自己那一列。说白了:attention 在 token 之间搬运信息;FFN 在 channel 之间搬运信息,一次只处理一个 token。
第二步故意花钱多。一个 block 的 12d² 个参数,拖动宽度,看它们怎么分:
FFN 拿走每个 block 三分之二的权重 —— 因为每个 block 的隐藏层比 d_model 宽 4 倍,这笔钱要付两次,升维一次、降维一次。试试:具体数量缩小了,但这个比例从来没变过。它是个比值,不是个数量。
block 还要选一件事:归一化放在哪里。切换一下,看梯度自己回到 embedding 的路:
在 post-norm 下 —— 也就是 2017 年最早的设计 —— 残差路径本身在每一个 block 都要穿过一次归一化,N 个连起来一次不落。在 pre-norm 下,归一化只碰支路,从不碰 stream,梯度穿过的次数就是 0。就这一个选择,基本上就是 GPT-2、Llama 以及几乎所有现代模型能训过百层还稳的原因。 §03 离开这座 stack,问一句:最后从另一头出来的是什么。
从 logits 到一个 token
stack 的尽头是 logits,每个词表条目一个分数。把它变成一个选定的词,是它自己的一台小机器 —— 这一页里第一件 stack 前面没有造过的事。
取故事里的一个点:prompt 是「the cat sat on the ___」,模型给 5 个候选打了分 —— mat、rug、floor、sofa、roof。一步步看每一个的原始分数,在任何东西改变它之前:
注意光看原始数字「mat」就已经领先,4.0 甩开其余几个。 softmax 把这一串分数变成加起来等于 1 的概率;真正改变这个分布形状的 —— 不是改变谁领先 —— 是温度。目前领先的候选会随着你拖滑块重新画出来。静止在 T = 1 时,「mat」已经占了 68.7%:
看 T 往 1 以上爬时,这些柱子怎么被拉平;看 T 往 0 掉时,「mat」怎么几乎吞掉所有概率。到时它已经拿走了几乎全部概率 —— 这就是 argmax 在极限下的样子:永远只取那个最高的 logit,每次都一样。
argmax 是切掉尾巴的一种办法;只留前 k 个候选是另一种。把 k 往下拖,看被切掉的概率怎么变大:
时这就是 argmax —— 一个候选,零个备选。注意 k 从不随分布的形状调整:不管模型有多确信,它都切到同一个固定数量。
top-p 是自适应版本 —— 只留累计概率刚好达到 p 的最小集合,分布越尖锐留的候选就越少,越平就留得越多:
时恰好留下两个候选 —— mat 和 rug 加起来刚过 90%, 其余的都没进来。把 p 往下推,这个集合能缩到只剩一个;把 p 推向 1,只要有一点概率的候选就全都留着。
贪心解码 —— 每一步都 argmax,一直这样 —— 有一个真实的失效模式:亲手把它逼出来看看:
一旦这个玩具模型的 argmax 路径绕回自己,贪心解码就会把那个 token重复到天荒地老 —— 它没有机制能发现这一点,也没有出路。在正确的时刻插入一次采样就能打破这个循环;生产系统里,这个现象有个名字,叫重复坍缩(repetition collapse)。 §04 拿走这一节刚选出的 token,问一句:把它喂回去的那一刻,会发生什么。
一次生成一个 token
一个只会给下一个词打分的模型,还不是生成器。把赢家喂回去,当它本来就是prompt 的一部分 —— 它就是了。
每过一遍 stack,序列就恰好长一个 token:整个跑一遍、采样、接上去、再来一遍。没有单独的「生成模式」—— 就是 §01 到 §03 已经造好的那同一次前向传播,只是每次拿更长一点的输入再调用一次。
prompt是给定的,从一开始就是灰的;每个生成出来的 token都会安定在它后面;滑块下面那一个是模型此刻正在决定的东西。一步一步走过去:
注意每个已经安定的 token都会变成下一个 token 的输入的一部分 —— 这正是「自回归」的全部定义:第 t 步的输出,是第 t+1 步输入的一份原料。
不小心的话,这要付出的代价是这样的。naive 的做法每一步都把已有的每个 token 重新跑一遍每一层;聪明一点的做法只处理新出现的那个。在同一步上对比一下:
在「naive」下,每往前一步整个前缀都要重新点亮一遍 —— 用 n 个 token 的工作量,换来第 n+1 个 token。在「cached」下,只有最新那格会亮。
把这个差距摊到一次真正的生成上 —— 玩具规模,4 层,d = 64 —— 看两条总量怎么越拉越开:
第一个 token 时两边的成本还很接近 —— 差 5 倍。到时,naive 的总量已经是cached 总量的十倍还多,而且差距还在拉大:naive 的成本随已生成量的平方增长,cached 的成本只线性增长。
把这个放大到 GPT-2 small 自己的 12 层、768 维,从一个 20 token 的 prompt开始:
拖到,比值落在 69 倍开外。这不是舍入误差 —— 这正是生产系统从来不会真的重跑前缀的全部原因。它们改存的是什么、存了多少,就是 §05。
cache 到底记住了什么
「cached」存的到底是什么:attention 已经算出来的每一个 key 向量和 value 向量 —— 每个 token、每个 head、每一层各一对,存下来就再也不碰。
attention 是拿一个新 token 的 query,去跟它前面所有的 key 打分。那些旧的 key 和 value 一旦写下来就不会变 —— 产生它们的那个 block 再也不会在它们身上重新跑一遍 —— 所以 cache 只是一个老老实实存它们的地方。
每生成一个 token,恰好追加一对。看填满的槽位怎么累积:
这就是这一整节都靠着的那条操作不变式:生成完 token t 之后,cache 恰好存着 t 对;生成 token t+1 只会再追加一对 ——已经在那儿的那些,既不会重算,也不会重写。
把这笔账放大到 GPT-3 自己发表过的形状 —— 96 层,d_model = 12,288 ——它就不再是免费的了。把上下文拖出去:
在 GPT-3 自己的 2,048 token 上下文下,光是一条序列的 cache就到了 9 GiB, fp16 精度 —— 还没加载一个权重呢。退回到,正好是它的四分之一:公式是线性的,每个数 2 字节,一个 key 向量加一个 value 向量,每层一对,再乘以窗口里到底有多少个 token。
cache 通常是个固定大小的缓冲区,不是个无限长的列表。把生成推过它自己的容量,看一个 naive 的环形缓冲区会做什么:
过了第 7 个槽位,写入就绕回来,落在一个还有效的更早 token 需要的槽位上。如果没有东西守住缓冲区的边界,这个失败是悄无声息的:没有报错,没有崩溃,只是一个更老 token 的 key被更新的悄悄换掉,从那以后基于它算出来的每个分数都是错的。
每个被缓存的 key 还得带一样东西同行:它自己坐在哪个位置。切换一个真实存在过的实现 bug,看它怎么跑偏:
把这个搞错 —— 每个新 token 都写 position 0,而不是 cache 自己当前的长度 ——位置编码从 §01 开始这一整页都在假设它成立,此刻就对不上了:模型把每个新 token 都当成第一个来打分,输出悄悄劣化,一个异常都不会抛出来。 §06 给这一切实际跑起来到底要花多少钱,安上一个数字。
这笔账
两个不同的问题,都值得给一个真实的数字:跑一遍到底花多少钱,内存到底流到哪儿去了。
每一步都要把整个权重矩阵从内存里读一遍,然后花在那一遍里不管有多少个 token 上。这个比值 —— 每读一字节权重能换来多少 FLOPs —— 决定了这一步到底是被芯片的算力卡住,还是被内存喂数据的速度卡住。
prefill 一次把 prompt 里所有 token 都打完分,所以一次权重读取服务了它们全部; decode 每次只服务一个。看这两根柱子随 prompt 变长怎么拉开:
prefill 的强度随 prompt 一起爬升 —— 推到,每读一字节就能换来 16 个 FLOPs,妥妥的算力受限。decode 的强度不管模型多大,都恰好钉在 1 附近 —— 这是内存带宽受限,再多算力也救不了它。
再看另一个问题。在 GPT-3 自己的规模上,把固定的权重和一条序列自己的 cache放在一起称一称:
就算是满打满算的 2,048 token 上下文,一条序列的 cache摆在 325 GiB 的权重旁边也就是一条细缝。一个请求便宜。账本变样是从第二个请求开始的。
权重只付一次钱,被所有在跑的请求共用;cache 不是 —— 它是按序列算的。拖一下并发序列数:
超过同时挂着,它们加起来的 cache就比服务它们所有人的权重还要占更多内存 —— 正是这个压力,让 KV cache 的内存管理、而不是原始 FLOPs, 成了生产级服务系统真正围着转的东西。
最后一个数字,把这条线绕回 §01 和 §02:GPT-2 small 自己的1.24 亿参数里:
大约三分之二住在那 12 个 block 里 —— 按 §02 的说法,大头是FFN —— 剩下的就是这一整页从查表开始就在用的那张 embedding 表。§07 把这个循环收进可以运行的代码里,把这些成本收进一张表里,再指出接下来该读什么。
速查
整个循环收进 9 行代码,成本收进一张表,6 个阶段再看一遍,加上那个把它们变成生成器的箭头。
# grow the sequence one token at a time
cache = None
while len(tokens) < max_len:
x = tokens[-1:] if cache else tokens
logits, cache = model(x, cache)
probs = softmax(logits[-1] / temperature)
next_id = sample(probs, top_k, top_p)
tokens.append(next_id)
if next_id == EOS_ID: break这个循环里的每一行,都是这一页某一节已经给它安上过数字的东西 —— 一行一行走一遍,看看是哪一节:
这四行加起来到底要花多少钱:
为什么这个界是紧的
天花板很少是原始 FLOPs。prefill 是算力受限的;decode 不管模型多大都钉在大约每读一字节一个 FLOPs 附近,这从构造上就是内存带宽受限 —— 而超过几十条并发序列之后,KV cache 占的内存就会超过服务它们所有人的权重。
变体
把循环重新打开: