多生成一个 token 时,什么被重新计算了
一句话总结
自回归生成一次只产生一个令牌。没有缓存时,每多生成一个令牌,就要把前面整段上下文的 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 的矩阵。
没有缓存时生成一个令牌:
- 对前面所有位置生成 q·k·v——3 乘 n 乘 d 乘 d 次
- 最后一个 query 查看 n 个 key——n 乘 d 次
- 用这些权重混合 n 个 value——n 乘 d 次
使用缓存时生成一个令牌:
- 只对新令牌一个生成 q·k·v——3 乘 d 乘 d 次
- query 查看 n 个 key——n 乘 d 次
- 用这些权重混合 n 个 value——n 乘 d 次
区别只在第一行。没有缓存的一边,投影成本随上下文长度成正比增长,而使用缓存的一边,这部分成本与长度无关,是固定的。注意力本身(后面两行)两边都随长度成正比增长——缓存并没有把它去掉。缓存去掉的是重复计算,而不是注意力。
本实验要亲手数出这些数,按长度做成表。不是照抄别人写下的倍数,而是写上自己数出来的数。
缓存成立的原因是因果掩码
为什么前面位置的 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 的注意力、无缓存运行的版本、使用缓存运行的版本。然后用同样的输入运行两个版本,用容许误差确认输出是否为相同的值,再按长度数出乘法次数,做成表。
最后两步是本实验的要点。四舍五入后放进缓存,两种方式的输出就不再落在容许误差之内——缓存不是近似,但写进缓存时不那么精确就是近似,这一点会用数字呈现出来。接着用你定的层数、头数、头维度和数据类型,计算缓存占用的元素个数和字节数,与乘法次数表一起留成记录。评分器会真正导入你的模块,每次用不同大小的输入调用函数,并另行计算值与乘法次数来对照。