数一数 KV 缓存省下的乘法
目标
通过亲手数乘法次数,确认在自回归生成中 KV 缓存省下了什么。做出数乘法的乘法函数,分别做出无缓存运行的版本和使用缓存运行的版本,用容许误差确认两个版本的输出是否为相同的值,再按上下文长度数出乘法次数,做成表。最后测量缓存若以较低精度保存,这种相同性会如何被破坏,并用自己定的设置计算缓存占用的元素个数和字节数,留成记录。
为什么重要
亲手算一次注意力,和在生成循环里调用这个注意力,是两回事。要生成长度为 100 的回答,就要调用模型 100 次,没有缓存时每次调用都要把前面整段上下文的 k·v 从头再算一遍。也就是说,第一个令牌的 k·v 生成一百次,丢掉九十九次。
这种浪费光读代码是看不出来的,因为注意力函数本身没有任何错误的地方。要让它显形,就得数。而且不能测时间——时间会因机器和负载而不同,但乘法次数在同样的输入下永远相同,并且会原样展示出它随长度如何增长。
本实验不调用实际模型。这个 Pod 的系统 Python 里没有 numpy(只在 /opt/onnx-lab/bin/python 里有),也没有 torch 和 transformers。只用标准库搭出同样的结构,只使用在这里测出来的数字。所以“实际模型会占用多少 GB”或“快多少倍”之类的话,这里不说。
评分器不会相信你写下的说明。它会真正导入你的模块,每次用不同大小的输入调用函数,并另行计算值与乘法次数来对照。大小每次运行都会变,所以无法把值背下来填进去。
步骤
- 在 /root/work/tf-kv/kv.py 中创建
CONFIG以及reset_muls()、muls()、mul(a, b)、dot(u, v)。mul在做乘法的同时把计数加一,dot只通过这个mul来做乘法。 - 增加
project(x, W),让一个令牌通过一组权重。W是每行长度为len(x)的矩阵,返回的列表长度为len(W)。 - 增加
attend(q, K, V),让一个 query 查看已累积的整个 K·V。把分数除以√len(q),套上 softmax,再对 V 做加权平均。 - 增加
step_nocache(xs, Wq, Wk, Wv),不用缓存得出一个输出。对前面所有令牌重新生成 q·k·v,再用最后一个 query 做注意力。 - 增加
new_cache()、store(cache, k, v, digits=None)、step_cached(x, Wq, Wk, Wv, cache, digits=None),使用缓存得出同样的输出。只计算新令牌一个,并在 K·V 中追加一行。 - 增加
ATOL、RTOL、close(got, want)、compare(xs, Wq, Wk, Wv, digits=None),比较两种方式的输出是否为相同的值。不用等号比较。 - 增加
mul_table(xs, Wq, Wk, Wv),数出上下文长度从 1 到len(xs)时两种方式的乘法次数。返回的值是(문맥길이, 캐시없음, 캐시있음)(占位符依次为上下文长度、无缓存、有缓存)组成的列表。 - 用 12 个令牌做出表,计算缓存占用的内存,记录到 /root/work/tf-kv/kv_report.json 和 /root/work/tf-kv/kv_report.md 中。
参考
- 执行契约:评分器会把
/root/work/tf-kv/kv.py当作 Python 模块导入,直接使用CONFIG、reset_muls、muls、mul、dot、project、attend、step_nocache、new_cache、store、step_cached、ATOL、RTOL、close、compare、mul_table。它不会作为脚本运行,所以可以没有if __name__ == "__main__"。 CONFIG是带有{"layers": ..., "heads": ..., "head_dim": ..., "dtype_bytes": ...}四个键的字典。值由你来定。范围是:层数 2 到 8,头数 2 到 8,头维度为 4 到 32 的偶数,数据类型字节数为 1、2、4 之一。不需要模仿实际模型的数——第 8 步的内存计算只用你定的这些值。mul(a, b)返回a * b,同时把模块内的计数器加 1。reset_muls()把计数器归 0,muls()返回当前值。所有发生乘法的地方都必须经过mul,数才对得上。dot(u, v)是点积。乘法发生len(u)次。加法、除法、指数不计入。project(x, W)的乘法是len(W)乘len(x)次。W的行生成输出的一格——把行和列混用,值和次数都会错位。attend(q, K, V)的乘法,分数一侧是len(K)乘len(q)次,加权和一侧是len(V)乘len(V[0])次。除法和指数不是乘法,所以不计入。softmax 要先减去最大值再取指数。step_nocache(xs, Wq, Wk, Wv)对xs的所有令牌生成 q·k·v,并用最后一个 query 查看整个 K·V。要用的 query 明明只有一个却全部生成,这就是没有缓存的样子。new_cache()返回{"K": [], "V": []}。store(cache, k, v, digits=None)在 K 和 V 中各追加一行——不覆盖。给出digits时,把每个分量四舍五入到该小数位数再放入。step_cached(x, Wq, Wk, Wv, cache, digits=None)只对新令牌一个做投影,用store放入缓存之后,再用这个 query 查看整个缓存。放入之前就做注意力的话,新令牌就看不到它自己。- 令
ATOL = 1e-9、RTOL = 1e-6,close(got, want)是abs(got - want) <= ATOL + RTOL * abs(want)。请用注释写下为什么是这个范围。Python 的 float 是 IEEE 754 双精度,有效数字约 15 位,这个规模的点积和 softmax 即使只改变相加的顺序,相对误差也只停留在 1e-12 的水平。 compare(xs, Wq, Wk, Wv, digits=None)从长度 1 到len(xs)每次增加一格地运行两种方式,返回{"lengths": [...], "max_gap": 실수, "all_close": 참거짓, "last_out": [...]}(占位符依次为实数与布尔值)。max_gap是各分量之差的绝对值中最大的一个,all_close表示所有分量是否都通过了close。mul_table(xs, Wq, Wk, Wv)在每个长度测量之前把计数器归零。缓存一侧必须持续沿用同一份缓存——每行都重新做一份,那就不是缓存。- 第 8 步令
d = CONFIG["head_dim"],用 12 个令牌和三个d乘d的权重来测量。round_digits用 2。 - 第 8 步的令牌向量和权重用
((i * 7 + j * 5 + salt * 3) % 13 - 6) / 7.0来生成。i是行号,j是列号,salt对令牌向量是 1,对 Wq 是 2,对 Wk 是 4,对 Wv 是 6。乘法次数和内存只取决于大小而不取决于值,所以不论用什么值,表都一样,但四舍五入的实验取决于值——除以 7 是为了不让值在小数点后第二位恰好除尽。如果只用恰好除尽的值,四舍五入到小数点后第二位值也不变,就看不出四舍五入的影响。 - 内存是
원소 수 = 2 × layers × heads × 길이 × head_dim(韩文,意为“元素个数 = 2 × layers × heads × 长度 × head_dim”)和바이트 = 원소 수 × dtype_bytes(韩文,意为“字节数 = 元素个数 × dtype_bytes”)。长度为 12。 - 这个 Pod 没有互联网。
pip install无法使用,也没有 torch 和 transformers。numpy 只在/opt/onnx-lab/bin/python里,所以在系统 Python 中import numpy不可用。只用math就足够了。 - 不要测时间。用
time测出来的数会因机器和负载而不同,不能用于判定。本实验测的是乘法次数。 - 官方文档:Attention Is All You Need · Hugging Face — Cache strategies · Hugging Face — Text generation · Python — math
- 常见错误:不经过
mul而直接相乘,导致计数为 0;在project中把行和列混用;分数没有除以√d;对缓存做覆盖;放入缓存之前就做注意力;在mul_table中没有把计数器归零;每行都重新做一份缓存。
数乘法的乘法
在 /root/work/tf-kv/kv.py 中创建 CONFIG 以及 reset_muls()、muls()、mul(a, b)、dot(u, v)。CONFIG 是带有 layers、heads、head_dim、dtype_bytes 四个键的字典,值由你来定(层数 2 到 8,头数 2 到 8,头维度为 4 到 32 的偶数,数据类型字节数为 1、2、4 之一)。mul 在做乘法的同时把计数器加 1,dot 只通过这个 mul 来做乘法。
计数器只要是模块里的一个整数就行。要在函数里修改它,需要 global。如果 dot 像 sum(a * b for a, b in zip(u, v)) 那样直接相乘,就什么也数不到,所以一定要经过 mul。加法不计入——矩阵计算的值来自乘法一侧。不要放进测时间的代码。
对一个令牌做投影
增加 project(x, W)。W 是每行长度为 len(x) 的矩阵,一行生成输出的一格。返回的列表长度为 len(W),乘法发生 len(W) 乘 len(x) 次。
对每一行调用一次前面做好的 dot 就完了。可以写成一行。把行和列混用,在方阵里只是值错,在长方形矩阵里连长度也会错位,所以请先确认返回的列表长度是否为 len(W)。q、k、v 都用这一个函数来生成。
一个 query 查看整个缓存
增加 attend(q, K, V)。把 q 与每个 k 的点积除以 √len(q) 得到分数,减去最大值后取指数套上 softmax,再用这些权重对 V 做加权平均。乘法,分数一侧是 len(K) 乘 len(q) 次,加权和一侧是 len(V) 乘 len(V[0]) 次。
除法和指数不是乘法,所以不要用 mul 包起来——包起来的话次数就会错位。只有加权和的乘法才经过 mul。V 一行的长度可能与 q 的长度不同,所以输出列表要按 len(V[0]) 来设。如果 K 只有一行,权重就是单独的 1,输出应该与 V[0] 相同。
没有缓存——把前面的全部重新计算
增加 step_nocache(xs, Wq, Wk, Wv)。对 xs 的所有令牌生成 q·k·v,并用最后一个 query 查看整个 K·V。要用的 query 只有一个却全部生成,这就是没有缓存的样子。
三次列表推导式加一次 attend 就够了。不能因为只用最后一个 query,就只对最后一个令牌做投影——key 和 value 必须对前面所有令牌都有,而这一步的要点就是它们会被每次重新生成。乘法次数随上下文长度成正比增长。
使用缓存——只追加一行
增加 new_cache()、store(cache, k, v, digits=None)、step_cached(x, Wq, Wk, Wv, cache, digits=None)。new_cache() 返回 {"K": [], "V": []},store 在 K 和 V 中各追加一行(给出 digits 时按该小数位数四舍五入后放入),step_cached 只对新令牌一个做投影并放入缓存,之后再用这个 query 查看整个缓存。
不要覆盖,要 append。前面各行的值不会因新令牌的加入而改变——因为有因果掩码,每个位置只看自己之前的位置,这就是缓存成立的原因。如果在放入之前就做注意力,新令牌就看不到它自己,答案会与无缓存的版本不同。乘法次数方面,投影一侧与上下文长度无关,是固定的。
两种方式的输出是否为相同的值
增加 ATOL = 1e-9、RTOL = 1e-6、close(got, want)、compare(xs, Wq, Wk, Wv, digits=None)。close 是 abs(got - want) <= ATOL + RTOL * abs(want),并用注释写下为什么是这个范围。compare 从长度 1 开始每次增加一格地运行两种方式,返回 {"lengths": [...], "max_gap": 실수, "all_close": 참거짓, "last_out": [...]}(占位符依次为实数与布尔值)。
不要用等号来比较。两种方式碰巧按同样的顺序相加时,甚至可能连比特都相同,但只要相加的顺序稍有不同,最后一位就会晃动——所以正确的测试是容许误差。digits 原样传给 step_cached。四舍五入后放入,all_close 会变成假,这正是这一步想让你看到的。max_gap 是在所有长度、所有分量中最大的差。
按长度数乘法
增加 mul_table(xs, Wq, Wk, Wv)。数出上下文长度从 1 到 len(xs) 时两种方式的乘法次数,返回 (문맥길이, 캐시없음, 캐시있음)(占位符依次为上下文长度、无缓存、有缓存)组成的列表。每次测量之前把计数器归零,缓存一侧持续沿用同一份缓存。
调用两次 reset_muls()——测无缓存版本之前一次,测缓存版本之前一次。如果不归零,前面的数就会混进后面的数里。如果每行都重新做一份缓存,缓存一侧的数就会像无缓存一侧那样增长——那就不是缓存。看表可以发现,无缓存一侧随长度成正比增加,缓存一侧只有注意力那一份在增加。
把省下的和付出的一起写下来
令 d = CONFIG["head_dim"],用 12 个令牌和三个 d 乘 d 的权重运行 mul_table 和 compare。值用 ((i * 7 + j * 5 + salt * 3) % 13 - 6) / 7.0 来生成,salt 对令牌向量是 1,对 Wq 是 2,对 Wk 是 4,对 Wv 是 6。compare 不做四舍五入运行一次,再用 digits=2 运行一次。然后在 /root/work/tf-kv/kv_report.json 中写入 layers、heads、head_dim、dtype_bytes、tokens、table、total_nocache、total_cached、saved_muls、cache_elems、cache_bytes、round_digits、max_gap_exact、all_close_exact、max_gap_rounded、all_close_rounded,并在 /root/work/tf-kv/kv_report.md 中用 ## 무엇을 쟀나(韩文,意为“测量了什么”)、## 캐시 없이 하면 무엇을 다시 계산하나(韩文,意为“没有缓存时要重新计算什么”)、## 캐시가 먹는 메모리(韩文,意为“缓存占用的内存”)、## 두 방식의 출력이 같은가(韩文,意为“两种方式的输出是否相同”)四节来写。
数字不要手写,要用实际运行你的代码得到的值来填。total_nocache、total_cached 是把表中各格相加,saved_muls 是它们的差。cache_elems 是 2 × layers × heads × 12 × head_dim,cache_bytes 是再乘以 dtype_bytes 的值。all_close_exact 应为真,all_close_rounded 应为假——缓存本身不是近似,但写进缓存时不那么精确就是近似。生成值的式子里除以 7,是为了不让值在小数点后第二位恰好除尽。如果只用恰好除尽的值,四舍五入后值也不变,这个实验本身就无法成立。不要测时间。