TT Lab
开始
学习 学习路径 课程

Transformer — 手算一遍注意力

多生成一个 token 时,什么被重新计算了

在 TT Lab 中继续学习

一句话总结

自回归生成一次只产生一个令牌。没有缓存时,每多生成一个令牌,就要把前面整段上下文的 k·v 从头再算一遍。KV 缓存把这种重算去掉了——公式原封不动,只改变计算的次数。

为什么需要它

注意力本身,你应该已经亲手算过了。打分,套上 softmax,对 value 做加权平均。就三行。

问题出在在生成循环里调用这三行的时候。要生成长度为 100 的回答,就要调用模型 100 次。每调用一次,模型都会把“到目前为止的整段上下文”作为输入,并且对这段上下文的每个位置都生成 q·k·v。

可是第二次调用时生成的第一个令牌的 k·v,与第一次调用时生成的是完全相同的值。第三次也一样,第一百次也一样。同样的值生成一百次,丢掉九十九次。

这种浪费光读代码是看不出来的,因为注意力函数里没有任何错误。要让它显形,就得数次数。

为什么数乘法而不是测时间

“打开缓存后变快了”这种测量,什么也学不到。变快多少,取决于机器、负载、批次大小和内存带宽,在同一台机器上重测一次也会得到不同的数。

乘法次数则不同。同样的代码,同样的输入,数出来的永远是同一个数。而且这个数会原样展示出它随长度如何增长。所以本实验不测时间。取而代之的是做一个执行乘法的函数,在里面计数。

def mul(a, b):
    global _MULS
    _MULS += 1
    return a * b

这是老实的办法。让所有发生乘法的地方都经过这个函数,就能数出谁在哪里乘了几次。计算量不靠猜,靠数。

数出来的样子

只取一层一个头,模型维度记为 d,当前上下文长度记为 n。生成 q·k·v 的权重是三个 d 乘 d 的矩阵。

没有缓存时生成一个令牌:

使用缓存时生成一个令牌:

区别只在第一行。没有缓存的一边,投影成本随上下文长度成正比增长,而使用缓存的一边,这部分成本与长度无关,是固定的。注意力本身(后面两行)两边都随长度成正比增长——缓存并没有把它去掉。缓存去掉的是重复计算,而不是注意力。

数出生成五个令牌期间每个位置的 key·value 被生成了几次的图。没有缓存时每次都要把前面的位置全部重新生成,依次是 1、2、3、4、5 次,共 15 次;使用缓存时只生成新加入的一个位置,每次 1 次,共 5 次

本实验要亲手数出这些数,按长度做成表。不是照抄别人写下的倍数,而是写上自己数出来的数。

缓存成立的原因是因果掩码

为什么前面位置的 k·v 可以原样复用?新令牌接在后面了,凭什么确信前面位置的值不会变?

因为有因果掩码。Attention Is All You Need 的解码器规定每个位置只看自己之前的位置。第 3 个位置的 k 和 v 只由第 3 个位置的输入生成,第 4 个令牌接在后面,第 3 个位置的值没有理由改变。

如果没有掩码,所有位置互相都能看到,缓存就不成立。因为每次后面接上新内容,前面的表示都会改变。所以 KV 缓存是仅解码器结构的性质,而不是到处都能用的优化。

这里能得出一个重要的结论:缓存不是近似。答案不会变。只是不再重新生成同样的值而已。所以打开和关闭缓存时输出若不同,那不是缓存的性质,而是实现的缺陷。

缓存占用的内存

有省下的,就有付出的。缓存要占用内存。元素个数写成乘积。

원소 수 = 2 (K 와 V) × 층 수 × 헤드 수 × 문맥 길이 × 헤드 차원
바이트  = 원소 수 × 자료형 한 원소의 바이트 수

这里值得注意的是上下文长度是乘进去的。长度翻倍,缓存也翻倍。这是长上下文昂贵的原因之一,一台机器能同时处理多少个请求,通常也由这张表决定。

本实验把层数、头数、头维度设为你自己定的较小的值,只用这些值来计算。不使用实际模型的 GB 数字——因为这里没有测过。把没测过的数字抄进来的那一刻,这份记录就失去了依据。更换数据类型时只有最后一项会变,长度是乘进去的,这两点在自己做出的表里能直接看到。

在现场相遇的样子

第一,每次都重新发送很长的提示词,用不上缓存。接着对话时如果把前面的内容整个重发一遍,服务器一侧就没有依据继续使用缓存。Hugging Face 的 Cache strategies 文档之所以专门讲解在两次调用之间亲手携带缓存对象的方法,原因就在这里。

第二,并发请求数的上限来自内存,而不是计算。缓存按请求分别占用,并随长度成正比增长。所以几个很长的对话,比几十个很短的请求占用更多空间。

第三,用较低的精度保存缓存,输出会出现细微的差别。缓存本身不是近似,但写进缓存时不那么精确是近似。本实验会在放入缓存之前先四舍五入,亲自确认这个差别是否超出容许误差。

第四,第一个令牌和后面令牌的性质不同。整段读入提示词的第一次计算,缓存是空的,没什么可省。缓存带来收益是从第二个令牌开始的。Hugging Face 的 Text generation 文档也把生成分成这两个阶段来讲解。

第五,使用缓存时混用批次,位置会错位。缓存的行顺序就是令牌顺序。把请求合并又拆开,如果把行接错了,不会报错,只会让输出变得奇怪。

实际工作中真正重要的事

下一项实验要做什么

把 /root/work/tf-kv/kv.py 一步一步做大。只使用标准库——这个 Pod 的系统 Python 里没有 numpy(只在 /opt/onnx-lab/bin/python 里有),也没有 torch 和 transformers。所以这里出现的数字,全部是你自己写的代码数出来的。

从数乘法的乘法函数和点积开始,依次做出:对一个令牌做投影的函数、一个 query 查看已累积的 K·V 的注意力、无缓存运行的版本、使用缓存运行的版本。然后用同样的输入运行两个版本,用容许误差确认输出是否为相同的值,再按长度数出乘法次数,做成表。

最后两步是本实验的要点。四舍五入后放进缓存,两种方式的输出就不再落在容许误差之内——缓存不是近似,但写进缓存时不那么精确就是近似,这一点会用数字呈现出来。接着用你定的层数、头数、头维度和数据类型,计算缓存占用的元素个数和字节数,与乘法次数表一起留成记录。评分器会真正导入你的模块,每次用不同大小的输入调用函数,并另行计算值与乘法次数来对照。