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

MiniMind — 亲手从头到尾训练一个小型语言模型

用 MiniMind 的 generate 测量 KV 缓存与采样

在 TT Lab 中继续学习

目标

用 MiniMind 的 generate,确认有 KV 缓存和没有 KV 缓存时贪心解码是否生成相同的令牌,让进入 forward 的令牌数与公式对上,再测量时间。让缓存的实际字节数与公式对上,并用 SFT 模型数出温度和 top-p 改变了什么。

为什么重要

生成的成本由缓存决定。没有缓存,每生成一个新令牌都要把迄今为止的全部内容重新计算;有缓存,只计算一个令牌。这个差距随长度越长,以平方的方式拉开。代价是缓存要占内存,而它的大小可以由配置(层数、KV 头、head_dim)准确算出——在服务中,一张卡能接待多少人,就取决于这个数字。 取令牌的方法与缓存无关。缓存是“用更低的成本得到相同的结果”,采样是“得到什么结果”。把两者混在一起想,就会做出“打开缓存之后答案变了”这样的错误诊断。

步骤

  1. 用基准预训练模型,把验证文档最前面的 6 篇连起来截到 200 个令牌作为提示(前面加 bos),在其后分别带缓存和不带缓存用贪心解码生成 64 个令牌,把两者是否相同写入 /root/mm/infer/same.json。
  2. 数出这两次生成中进入 forward 的令牌数,写入 /root/mm/infer/count.json。
  3. 把两种方式的生成时间(预热一次后测三次取中间值)写入 /root/mm/infer/time.json。
  4. 测出放入 100 个令牌之后的 KV 缓存大小,与公式对上,写入 /root/mm/infer/kvbytes.json。
  5. 用 SFT 模型对“Garam 村的特产是什么?”分别以温度 0.3、1.0、1.5 各回答 30 次,把不同回答的个数写入 /root/mm/infer/temp.json。
  6. 对预训练模型放入 [bos] 민수는(韩文,意为“Minsu 这个人”)之后的下一个令牌分布,按 MiniMind 的 top-p 规则,把 p=0.5、0.9、0.99 时留下来的候选数写入 /root/mm/infer/topp.json。
  7. 在 /root/mm/infer/report.md 中写 ## KV 캐시、## 온도、## top-p 三节(三个标题为韩文,依次意为“KV 缓存”“温度”“top-p”),并放入第 3 步的 speedup 和第 4 步的缓存字节数。

参考

不管有没有缓存,回答都相同

编写脚本 /root/mm/infer/cache.py,加载 /opt/mm/ref/pretrain.pth,把 /opt/mm/data/pretrain_val.jsonl 最前面 6 篇文档的 text 用空格连起来切成令牌,取前 200 个并在前面加上 bos(1) 作为提示,在其后用 64 个令牌的贪心解码(do_sample=False, top_k=0, top_p=1.0, eos_token_id=None)分别以 use_cache=True、False 生成两次。把新令牌是否相同(same_ids)和带缓存一侧的新令牌(ids),连同 prompt_len、new_tokens 一起写入 /root/mm/infer/same.json。

多亏有因果掩码,前面令牌的 K、V 不会因为后面接上了令牌而改变。所以取出来用,结果也相同。如果这里不同,那不是缓存的问题,而是条件(采样、随机数)不同。

数出进入 forward 的令牌

把第 1 步两次生成中进入 model.forward 的 input_ids 的长度全部加起来,以 prompt_len、new_tokens、fed_with_cache、fed_without_cache 写入 /root/mm/infer/count.json。

有缓存的话,第一步放入整个提示,之后每次放一个令牌。没有的话,每一步都把迄今为止的全部内容重新放入。评分器会把你写的两个值与用这两个公式算出的值对照。

从时间上看

把带缓存和不带缓存的 64 个令牌生成各预热一次后分别测三次,取中间值,以 with_cache_s、without_cache_s、speedup(无缓存÷有缓存)写入 /root/mm/infer/time.json。

用 time.perf_counter() 来测。第一次运行因为内存分配和准备工作而变慢,所以排除。按令牌数相差 50 倍以上,时间却没有相差这么多——因为在小模型上,每一步的固定开销(Python、内核调用)很大。

缓存有多少字节

对基准预训练模型放入 100 个随机令牌(torch.randint(3, 어휘, (1, 100), generator=시드 0),占位符依次为词表大小与随机种子),用 use_cache=True,把返回的 past_key_values 所有张量的字节总和(measured_bytes)、用公式算出的值(formula_bytes = 2 × 层数 × 令牌数 × KV 头数 × head_dim × 4)、第一层 K 的形状(k_shape),连同 tokens 一起写入 /root/mm/infer/kvbytes.json。

MiniMind 的 Attention 在 repeat_kv 之前 生成 past_kv = (xk, xv)。所以 K 的形状是(批次,长度,KV 头 2,32),GQA 减少了多少,缓存也就小多少。

把温度调高会怎样

对 SFT 模型(/opt/mm/ref/sft.pth),以聊天格式(mmkit.chat_text(…, add_generation_prompt=True))放入“Garam 村的特产是什么?”,每个温度 0.3、1.0、1.5 都先设 torch.manual_seed(0),再以 do_sample=True, top_k=0, top_p=1.0, max_new_tokens=24, num_return_sequences=30 生成,把截到 <|im_end|> 之前的不同回答的个数,以 {"0.3": {"distinct": n}, "1.0": …, "1.5": …} 写入 /root/mm/infer/temp.json。

每个温度都要从相同的随机数出发,才能只看温度的效果。在温度 1.5 下,三十个回答几乎全都不同,还会出现字符损坏的回答——一旦选中了罕见的令牌,小模型就回不到正轨。

top-p 留下的候选

对基准预训练模型放入 [bos] + '민수는'(韩文,意为“Minsu 这个人”)之后最后一个位置的 logits,按 MiniMind generate 的 top-p 规则(从按概率降序的累积和超过 p 的位置起截断,但错开一位,保留第一次超过 p 的令牌,并且最上面的一个总是保留),把 p=0.5、0.9、0.99 时留下来的候选数,以 {"0.5": n, "0.9": n, "0.99": n} 写入 /root/mm/infer/topp.json。

把 /opt/minimind/model/model_minimind.py 的 generate 里 top_p 那三行原样搬过来就行。在“Minsu”这个主语之后有动词、地点等多条分支,所以会留下不少候选。请与分布很尖的位置(例如“村里的特产是”之后)比较。

推理成本与取令牌的记录

在 /root/mm/infer/report.md 中写 ## KV 캐시、## 온도、## top-p 三节(三个标题为韩文,依次意为“KV 缓存”“温度”“top-p”),并以数字放入第 3 步的 speedup 和第 4 步的 measured_bytes。

把缓存在令牌数上是几倍、在时间上是几倍并排写下来,再加一行说明这个差别的原因。