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

Transformer — 手算一遍注意力

数一数 KV 缓存省下的乘法

在 TT Lab 中继续学习

目标

通过亲手数乘法次数,确认在自回归生成中 KV 缓存省下了什么。做出数乘法的乘法函数,分别做出无缓存运行的版本和使用缓存运行的版本,用容许误差确认两个版本的输出是否为相同的值,再按上下文长度数出乘法次数,做成表。最后测量缓存若以较低精度保存,这种相同性会如何被破坏,并用自己定的设置计算缓存占用的元素个数和字节数,留成记录。

为什么重要

亲手算一次注意力,和在生成循环里调用这个注意力,是两回事。要生成长度为 100 的回答,就要调用模型 100 次,没有缓存时每次调用都要把前面整段上下文的 k·v 从头再算一遍。也就是说,第一个令牌的 k·v 生成一百次,丢掉九十九次。 这种浪费光读代码是看不出来的,因为注意力函数本身没有任何错误的地方。要让它显形,就得数。而且不能测时间——时间会因机器和负载而不同,但乘法次数在同样的输入下永远相同,并且会原样展示出它随长度如何增长。 本实验不调用实际模型。这个 Pod 的系统 Python 里没有 numpy(只在 /opt/onnx-lab/bin/python 里有),也没有 torch 和 transformers。只用标准库搭出同样的结构,只使用在这里测出来的数字。所以“实际模型会占用多少 GB”或“快多少倍”之类的话,这里不说。 评分器不会相信你写下的说明。它会真正导入你的模块,每次用不同大小的输入调用函数,并另行计算值与乘法次数来对照。大小每次运行都会变,所以无法把值背下来填进去。

步骤

  1. 在 /root/work/tf-kv/kv.py 中创建 CONFIG 以及 reset_muls()、muls()、mul(a, b)、dot(u, v)。mul 在做乘法的同时把计数加一,dot 只通过这个 mul 来做乘法。
  2. 增加 project(x, W),让一个令牌通过一组权重。W 是每行长度为 len(x) 的矩阵,返回的列表长度为 len(W)。
  3. 增加 attend(q, K, V),让一个 query 查看已累积的整个 K·V。把分数除以 √len(q),套上 softmax,再对 V 做加权平均。
  4. 增加 step_nocache(xs, Wq, Wk, Wv),不用缓存得出一个输出。对前面所有令牌重新生成 q·k·v,再用最后一个 query 做注意力。
  5. 增加 new_cache()、store(cache, k, v, digits=None)、step_cached(x, Wq, Wk, Wv, cache, digits=None),使用缓存得出同样的输出。只计算新令牌一个,并在 K·V 中追加一行。
  6. 增加 ATOL、RTOL、close(got, want)、compare(xs, Wq, Wk, Wv, digits=None),比较两种方式的输出是否为相同的值。不用等号比较。
  7. 增加 mul_table(xs, Wq, Wk, Wv),数出上下文长度从 1 到 len(xs) 时两种方式的乘法次数。返回的值是 (문맥길이, 캐시없음, 캐시있음)(占位符依次为上下文长度、无缓存、有缓存)组成的列表。
  8. 用 12 个令牌做出表,计算缓存占用的内存,记录到 /root/work/tf-kv/kv_report.json 和 /root/work/tf-kv/kv_report.md 中。

参考

数乘法的乘法

在 /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,是为了不让值在小数点后第二位恰好除尽。如果只用恰好除尽的值,四舍五入后值也不变,这个实验本身就无法成立。不要测时间。