""" KV 缓存:不存和存,到底差多少?在你自己的电脑上量一遍。 这不是模拟,是真的跑一个极小的注意力模型:一串固定随机种子生成的输入向量, 一个位置一个位置往下算,两种执行方式各跑一遍,量真实耗时、真实乘加次数和真实缓存占用。 先说清楚它**不是**什么:这里没有词表、没有分词器,也没有文本解码。 输入默认是 40 个随机向量,输出也是向量,屏幕上不会出现任何一个汉字。 所以它处理的是「位置」——在真实模型里一个位置对应一个 token,不是一个汉字。 「缓存占用」那一列同样不是它看上去的那个东西:它数的是缓存里存了多少个数, 不是内存读数。Python 的 list 与 float 每个元素实际要几十字节,而第 2 级旋钮里 「每数 2 字节」是 fp16 的逻辑张量存储估计——两者都不是这个进程真实占了多少内存, 更不是显存:这一趟是纯 CPU 本机运行,没有测量任何 GPU 显存。 只需要 python3,不用装任何东西: python3 kv_cache.py # 默认就跑这个,什么参数都不用给 python3 kv_cache.py --json # 同一份结果的机器可读版本 python3 kv_cache.py --help # 想换配置(种子、位置数、维度、层数、轮数)时再看 两条路径必须是同一个计算的两种执行方式,只差「要不要重算」。 这一点由 tests/demo/test_kv_cache.py 逐前缀、多层数、多维度、多种子地验证。 """ import argparse import datetime import hashlib import json import math import os import platform import random import statistics import time import unicodedata from pathlib import Path DIM = 32 LAYERS = 2 N_TOKENS = 40 DEFAULT_SEED = 0 # §3.6 建议的基准默认值:预热 2 轮、测量 10 轮。默认配置下整个脚本约 3~4 秒跑完。 DEFAULT_WARMUP = 2 DEFAULT_REPEATS = 10 # 机器可读记录的版本号。字段改了就要升版本,别让读它的人猜。 REPORT_SCHEMA = "ai-museum/kv-cache-benchmark/2" # §3.5:表里最后一列叫「缓存占用」,值是缓存里存了多少个数,**不是内存读数**。 # 这四句同时进机器可读记录和可读摘要,因为一个只看默认输出的读者最容易在这里 # 把「个数」读成「显存」。 CACHE_OCCUPANCY_NOTES = ( "最后一列数的是「缓存里存了多少个数」,不是内存占用", "Python 的 list 与 float 每个元素实际要几十字节,不等于「元素数 × 2 字节」", "「每数 2 字节」是 fp16 的逻辑张量存储估计,属于第 2 级旋钮那笔账", "纯 CPU 本机运行,没有测量任何 GPU 显存", ) SCRIPT_PATH = Path(__file__).resolve() def matvec(w, x): """一个向量过一次线性层。w 是 DIM×DIM。""" return [sum(wi * xi for wi, xi in zip(row, x)) for row in w] def softmax(xs): m = max(xs) es = [math.exp(v - m) for v in xs] s = sum(es) return [e / s for e in es] def attend(q, keys, values, cnt=None): """标准注意力:q 和每个 key 打分,softmax 后对 value 加权求和。""" scale = 1.0 / math.sqrt(len(q)) scores = [sum(a * b for a, b in zip(q, k)) * scale for k in keys] w = softmax(scores) out = [0.0] * len(q) for wi, v in zip(w, values): for j, vj in enumerate(v): out[j] += wi * vj if cnt is not None: cnt.score_macs += len(keys) * len(q) cnt.value_aggregation_macs += len(values) * len(q) return out class Counter: """数一数真的做了多少次乘加(MAC),这个数字是精确的,不受机器快慢影响。 分四项记,因为它们的增长方式不一样,混成一个总数就看不出谁在涨: kv_projection_macs 算 K 和 V 的投影 q_projection_macs 算 Q 的投影 score_macs Q 和每个 K 打分 value_aggregation_macs 按权重把 V 加起来 统计范围仅限上面这四项的乘加。softmax 的 exp 和除法、那个 1/sqrt(d) 的缩放乘、 循环与内存分配都**不在内**,所以 `mults` 是「这四项的乘加次数」,不是「所有计算」。 一次 MAC 是一次乘法加一次累加,不要不加说明地当成 FLOPs。 """ __slots__ = ("kv_projection_macs", "q_projection_macs", "score_macs", "value_aggregation_macs") def __init__(self): self.kv_projection_macs = 0 self.q_projection_macs = 0 self.score_macs = 0 self.value_aggregation_macs = 0 @property def mults(self): return ( self.kv_projection_macs + self.q_projection_macs + self.score_macs + self.value_aggregation_macs ) def build_model(rnd, layers=LAYERS, dim=DIM): def mat(): s = dim**-0.5 return [[rnd.gauss(0, s) for _ in range(dim)] for _ in range(dim)] return [{"q": mat(), "k": mat(), "v": mat()} for _ in range(layers)] def build_tokens(rnd, count=N_TOKENS, dim=DIM): """输入序列:count 个随机向量。不是文字,也没有经过分词器。""" return [[rnd.gauss(0, 1) for _ in range(dim)] for _ in range(count)] def forward_all_layers(model, tokens, cnt=None): """不存:把整个前缀从头过一遍所有层,返回每一层每个位置的输出。 一层里的顺序很关键:先用**这一层所有位置的输入**算 Q/K/V, 位置 i 只看 0..i,等所有位置都更新完,才把这一层的输出交给下一层。 少了这一步,第二层的历史位置就还吃着原始输入,跟缓存那条路径对不上。 """ if not tokens: raise ValueError("至少要有一个位置才能算") dim = len(tokens[0]) layer_in = [list(t) for t in tokens] per_layer = [] for layer in model: keys = [matvec(layer["k"], t) for t in layer_in] values = [matvec(layer["v"], t) for t in layer_in] if cnt is not None: cnt.kv_projection_macs += 2 * len(layer_in) * dim * dim out = [] for i, t in enumerate(layer_in): q = matvec(layer["q"], t) if cnt is not None: cnt.q_projection_macs += dim * dim out.append(attend(q, keys[: i + 1], values[: i + 1], cnt)) per_layer.append(out) layer_in = out return per_layer def gen_no_cache(model, tokens, cnt=None): """不存:每写一个新字,整个前缀都要从头重算一遍,只取最后一个位置的结果。""" return forward_all_layers(model, tokens, cnt)[-1][-1] def new_cache(layers=LAYERS): """开一次新的生成就得开一份新缓存,上一次的缓存不能接着用。""" return [{"k": [], "v": []} for _ in range(layers)] def gen_with_cache(model, new_token, cache, cnt=None): """存:只算新字的 K/V,前面的直接从缓存里取。 它和上面那条路径算的是同一个东西:第 li 层缓存里存的,正是历史位置在第 li 层的 K/V,而它们来自这些位置在第 li-1 层的输出——和重算路径里那些位置的值一模一样。 """ x = list(new_token) dim = len(x) for li, layer in enumerate(model): k = matvec(layer["k"], x) v = matvec(layer["v"], x) if cnt is not None: cnt.kv_projection_macs += 2 * dim * dim cache[li]["k"].append(k) cache[li]["v"].append(v) q = matvec(layer["q"], x) if cnt is not None: cnt.q_projection_macs += dim * dim x = attend(q, cache[li]["k"], cache[li]["v"], cnt) return x # -------------------------------------------------------------------------- # 计时与记录(执行文档 §3.6) # # 计时区里只有「算这一步」这一件事。权重、输入向量、前缀切片、缓存分配都在外面; # 乘加计数也单独跑一趟——计数器每次 += 都是真实开销,混进计时区就成了在量计数器。 # # 一「轮」是把全部位置完整跑一遍。缓存那一路每轮换一份新缓存,不然第二轮会接着 # 上一轮的历史往下算,量的就不是同一件事了。先预热若干轮再测量,两路的先后顺序 # 逐轮交替,免得固定次序让先跑的那一路单方面吃亏或占便宜。 # # 报中位数和最小值~最大值。毫秒级测量本来就抖,不要求逐个采样点单调: # 「注意力访问量随长度增长」不等于「每个采样点都比前一个慢」。 # # 计时区精确到「一步」:`t0 = perf_counter()` 到 `perf_counter()` 之间只有那一次生成调用。 # 所以「一轮总计」是那一轮各步测量值之和,不是这一轮的墙钟时间——轮与轮之间、 # 步与步之间的开销都不在里面。这一点在 report 的 `method.fullRunIsSumOfSteps` 里写着。 # -------------------------------------------------------------------------- def checkpoints(n): """挑几个位置打印,别刷满一屏。n=40 时正好是 5、10、20、30、40。""" return sorted({max(1, n // 8), max(1, n // 4), max(1, n // 2), max(1, 3 * n // 4), n}) def _run_no_cache(model, prefixes): """完整跑一轮「不缓存」,返回每一步的耗时(秒)。""" out = [] for prefix in prefixes: t0 = time.perf_counter() gen_no_cache(model, prefix) out.append(time.perf_counter() - t0) return out def _run_cached(model, tokens, layers): """完整跑一轮「缓存」,从一份新缓存开始。返回每步耗时(秒)和结束时的缓存长度。""" cache = new_cache(layers) out = [] for tok in tokens: t0 = time.perf_counter() gen_with_cache(model, tok, cache) out.append(time.perf_counter() - t0) return out, len(cache[0]["k"]) def _count_pass(model, tokens, prefixes, layers, dim, marks): """单独一趟数乘加,顺便在采样点记下真实的缓存元素数。这一趟不计时。""" c_no = Counter() for prefix in prefixes: gen_no_cache(model, prefix, c_no) c_yes = Counter() cache = new_cache(layers) elements = {} for i, tok in enumerate(tokens, start=1): gen_with_cache(model, tok, cache, c_yes) if i in marks: elements[i] = sum(len(c["k"]) * dim + len(c["v"]) * dim for c in cache) return c_no, c_yes, elements def _breakdown(cnt): return { "kvProjection": cnt.kv_projection_macs, "qProjection": cnt.q_projection_macs, "score": cnt.score_macs, "valueAggregation": cnt.value_aggregation_macs, "total": cnt.mults, } def _stats(samples_seconds): """秒 → 毫秒:中位数、最小值、最大值,外加全部样本,让别人能自己复算。""" ms = [s * 1000 for s in samples_seconds] return {"median": statistics.median(ms), "min": min(ms), "max": max(ms), "samples": ms} def _environment(): """§3.6 要记的环境字段。没有 GPU 就一个 GPU 字段都不写,不伪造。""" return { "pythonVersion": platform.python_version(), "pythonImplementation": platform.python_implementation(), "platform": platform.platform(), "machine": platform.machine(), "processor": platform.processor() or "未提供(platform.processor() 在这台机器上返回空)", "logicalCores": os.cpu_count(), "date": datetime.date.today().isoformat(), } def build_report( seed=DEFAULT_SEED, tokens=N_TOKENS, dim=DIM, layers=LAYERS, repeats=DEFAULT_REPEATS, warmup=DEFAULT_WARMUP, ): """跑一次基准,返回一份既能打印也能序列化的记录。""" if min(tokens, dim, layers, repeats) < 1 or warmup < 0: raise ValueError("位置数、维度、层数、测量轮数至少是 1,预热轮数至少是 0") # ---- 计时区外:权重、输入、前缀切片、乘加计数 ---- rnd = random.Random(seed) model = build_model(rnd, layers, dim) seq = build_tokens(rnd, tokens, dim) prefixes = [seq[:i] for i in range(1, tokens + 1)] marks = checkpoints(tokens) c_no, c_yes, elements = _count_pass(model, seq, prefixes, layers, dim, set(marks)) # ---- 预热:让解释器、分支预测和内存分配先稳下来(顺序同样逐轮交替)---- for w in range(warmup): if w % 2 == 0: _run_no_cache(model, prefixes) _run_cached(model, seq, layers) else: _run_cached(model, seq, layers) _run_no_cache(model, prefixes) # ---- 测量:逐轮交替两路的先后顺序 ---- rounds, no_cache_rounds, cached_rounds = [], [], [] for r in range(repeats): no_first = r % 2 == 0 if no_first: a = _run_no_cache(model, prefixes) b, tail = _run_cached(model, seq, layers) else: b, tail = _run_cached(model, seq, layers) a = _run_no_cache(model, prefixes) no_cache_rounds.append(a) cached_rounds.append(b) rounds.append( { "order": "no-cache-first" if no_first else "cached-first", "noCacheTotalMs": sum(a) * 1000, "cachedTotalMs": sum(b) * 1000, "cachedCacheLengthAtEnd": tail, } ) cps = [] for pos in marks: a = _stats([row[pos - 1] for row in no_cache_rounds]) b = _stats([row[pos - 1] for row in cached_rounds]) cps.append( { "position": pos, "cacheElements": elements[pos], "noCacheMs": a, "cachedMs": b, "medianRatio": (a["median"] / b["median"]) if b["median"] else None, } ) full_no = _stats([sum(row) for row in no_cache_rounds]) full_yes = _stats([sum(row) for row in cached_rounds]) return { "schema": REPORT_SCHEMA, "config": { "seed": seed, "tokens": tokens, "dim": dim, "layers": layers, "repeats": repeats, "warmup": warmup, }, "method": { "statistic": "median", "spread": "min-max", "warmupRounds": warmup, "measuredRounds": repeats, "orderBalance": { "noCacheFirst": sum(1 for r in rounds if r["order"] == "no-cache-first"), "cachedFirst": sum(1 for r in rounds if r["order"] == "cached-first"), }, "roundDefinition": "一轮 = 把全部位置完整跑一遍", "fullRunIsSumOfSteps": True, "cacheResetPerRound": True, "excludedFromTiming": ["权重生成", "输入向量生成", "前缀切片", "缓存分配", "乘加计数"], "acceleratorUsed": False, "acceleratorNote": "纯 Python、单线程、CPU;没有用加速器,所以不报告型号与同步耗时", "monotonicity": "不要求逐个采样点单调;毫秒级测量有抖动,严格增长的是计数那几列", }, "environment": _environment(), "script": { "name": SCRIPT_PATH.name, "sha256": hashlib.sha256(SCRIPT_PATH.read_bytes()).hexdigest(), }, "counts": { "account": "账 B:Q/K/V 投影 + 注意力打分 + 加权求和,共四项", "unit": "MAC(一次乘法加一次累加),不要不加说明地当成 FLOPs", "notCounted": ["softmax 的 exp 与除法", "1/sqrt(d) 缩放乘", "循环与内存分配"], "noCache": _breakdown(c_no), "cached": _breakdown(c_yes), "cachedShareOfNoCache": c_yes.mults / c_no.mults, }, # §3.5:最后一列是元素个数,不是内存。这段口径必须跟着记录走, # 不能只写在网页上——下载下来单独跑的人只看得到这份输出。 "cacheOccupancy": { "unit": "元素个数", "measures": "缓存 list 里存了多少个数,等于 2 × 层数 × 维度 × 位置数", "notMeasured": ["进程内存占用", "GPU 显存"], # 用 list,不用那个常量元组:`--json` 转一圈回来是 list, # 直接放元组会让「JSON 与 report 相等」这条断言在类型上先挂掉。 "notes": list(CACHE_OCCUPANCY_NOTES), }, "timing": { "unit": "ms", "checkpoints": cps, "fullRun": {"noCacheMs": full_no, "cachedMs": full_yes}, "rounds": rounds, }, } def _wide(text): """东亚宽字符按两列算,好让表头和数字列在等宽字体里真的对齐。""" return sum(2 if unicodedata.east_asian_width(ch) in "WF" else 1 for ch in text) def _pad(text, width): """按显示宽度右对齐。`f"{s:>8}"` 数的是字符数,中文表头会对不齐。""" return " " * max(0, width - _wide(text)) + text def render_text(report): """可读摘要。和 `--json` 是同一份 report 的两种渲染,不是两次测量。""" cfg, m, env = report["config"], report["method"], report["environment"] counts, timing = report["counts"], report["timing"] occ = report["cacheOccupancy"] bal = m["orderBalance"] lines = [ "KV 缓存:同一个计算的两种执行方式,在这台机器上各跑一遍", f"配置 {cfg['layers']} 层 × 每层维度 {cfg['dim']} × {cfg['tokens']} 个位置 · 随机种子 {cfg['seed']}", " 一个位置就是一个输入向量。没有词表、没有分词器,屏幕上不会出现汉字。", "计时口径 权重与输入向量在计时区外一次生成;乘加计数单独跑一趟,也不进计时区", f" 一轮 = 完整跑完 {cfg['tokens']} 个位置;缓存那一路每轮换一份新缓存", f" 预热 {m['warmupRounds']} 轮,测量 {m['measuredRounds']} 轮,两路先后顺序逐轮交替" f"(不缓存先 {bal['noCacheFirst']} 轮 / 缓存先 {bal['cachedFirst']} 轮)", f" 下面报 {m['measuredRounds']} 轮的中位数,方括号里是这几轮的最小值与最大值", " 「一轮总计」是那一轮各步测量值之和,不含轮与轮之间的开销", f"缓存占用 {occ['notes'][0]}", *[f" {note}" for note in occ["notes"][1:]], f"运行环境 Python {env['pythonVersion']}({env['pythonImplementation']})· {env['platform']}" f" · {env['machine']} · {env['logicalCores']} 逻辑核", f"运行日期 {env['date']}", f"脚本哈希 sha256 {report['script']['sha256']}", f"加速器 {m['acceleratorNote']}", "", _pad("位置", 8) + " " + _pad("不缓存·本步 [最小, 最大]", 26) + " " + _pad("缓存·本步 [最小, 最大]", 23) + " " + _pad("中位数比", 7) + " " + _pad("缓存占用", 11), "-" * 86, ] for cp in timing["checkpoints"]: a, b = cp["noCacheMs"], cp["cachedMs"] ratio = f"{cp['medianRatio']:>6.1f}×" if cp["medianRatio"] is not None else _pad("-", 7) lines.append( f"{cp['position']:>8} {a['median']:>6.2f} ms [{a['min']:>6.2f}, {a['max']:>6.2f}]" f" {b['median']:>5.2f} ms [{b['min']:>5.2f}, {b['max']:>5.2f}]" f" {ratio} {cp['cacheElements']:>6,} 个数" ) full_a, full_b = timing["fullRun"]["noCacheMs"], timing["fullRun"]["cachedMs"] lines += [ "-" * 86, _pad("一轮总计", 8) + f" {full_a['median']:>6.2f} ms [{full_a['min']:>6.2f}, {full_a['max']:>6.2f}]" f" {full_b['median']:>5.2f} ms [{full_b['min']:>5.2f}, {full_b['max']:>5.2f}]" f" {full_a['median'] / full_b['median']:>6.1f}×", f"总乘加次数 不缓存 {counts['noCache']['total']:>10,} 缓存 {counts['cached']['total']:>10,}", f" 缓存这一路是不缓存的 {counts['cachedShareOfNoCache']:.1%}", " (只统计 Q/K/V 投影与注意力这四项的乘加,不含 softmax、缩放和循环开销;", f" 这个配置:{cfg['layers']} 层 × 维度 {cfg['dim']} × {cfg['tokens']} 个位置,换配置数字就会变)", "", "注意最后两列:缓存之后本步耗时仍然在涨,缓存占用也一直在涨。", "缓存去掉的是「重算前面所有位置」,去不掉「每一步都要看一遍前面所有位置」。", "毫秒是中位数,不要求逐个采样点单调;严格增长的是那几列精确计数。", "想要机器可读的同一份结果:python3 kv_cache.py --json", ] return "\n".join(lines) def render_json(report): return json.dumps(report, ensure_ascii=False, indent=2) def _positive(name): def parse(text): value = int(text) if value < 1: raise argparse.ArgumentTypeError(f"{name}至少是 1,收到 {value}") return value return parse def _non_negative(text): value = int(text) if value < 0: raise argparse.ArgumentTypeError(f"预热轮数不能是负数,收到 {value}") return value def main(argv=None): parser = argparse.ArgumentParser( prog="kv_cache.py", description="KV 缓存:不存和存,两种执行方式在你自己的电脑上各跑一遍。", epilog="不带任何参数直接跑就行。下面这些是想换配置或要机器可读记录时才用的。", ) parser.add_argument("--json", action="store_true", help="输出机器可读 JSON(和可读摘要是同一份结果)") parser.add_argument("--seed", type=int, default=DEFAULT_SEED, help=f"随机种子,默认 {DEFAULT_SEED}") parser.add_argument("--tokens", type=_positive("位置数"), default=N_TOKENS, help=f"位置数,默认 {N_TOKENS}") parser.add_argument("--dim", type=_positive("维度"), default=DIM, help=f"每层维度,默认 {DIM}") parser.add_argument("--layers", type=_positive("层数"), default=LAYERS, help=f"层数,默认 {LAYERS}") parser.add_argument("--repeats", type=_positive("测量轮数"), default=DEFAULT_REPEATS, help=f"测量轮数,默认 {DEFAULT_REPEATS}") parser.add_argument("--warmup", type=_non_negative, default=DEFAULT_WARMUP, help=f"预热轮数,默认 {DEFAULT_WARMUP}") args = parser.parse_args(argv) report = build_report( seed=args.seed, tokens=args.tokens, dim=args.dim, layers=args.layers, repeats=args.repeats, warmup=args.warmup, ) print(render_json(report) if args.json else render_text(report)) if __name__ == "__main__": main()