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

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

KV 缓存就是记住无需重算的东西

在 TT Lab 中继续学习

一句话总结

生成就是每取出一个令牌,就把模型运行一次。没有缓存,每次都要把迄今为止的 全部内容 重新放入计算;有了 KV 缓存,就只放入新令牌 一个,前面令牌的 K、V 直接取出来用。结果连一个令牌都不会不同,只是计算变少了。而取令牌的方法(贪心解码、温度、top-p)则与此无关,决定的是“选哪个令牌”。这个模块会用 MiniMind 的 generate 把这两件事都用数字量出来。

为什么需要它

在 Transformer 课程里,我们用手算过 KV 缓存的原理和采样公式。这里要看的是,在 真正训练出来的模型 和 真实的生成循环 里,它们是什么样子。服务成本大部分出在生成阶段,而这笔成本的形状由缓存决定。而且同一个模型,用贪心解码总是给出同样的回答,把温度调高每次则给出不同的回答,这种差别究竟是“创造性”还是“胡说八道”,在小模型上表现得尤其鲜明。

工作原理

MiniMind 的 generate 是这样运转的。

for _ in range(max_new_tokens):
    past_len = past_key_values[0][0].shape[1] if past_key_values else 0
    outputs = self.forward(input_ids[:, past_len:], past_key_values=past_key_values, use_cache=use_cache)
    logits = outputs.logits[:, -1, :] / temperature
    ... top_k · top_p 로 자르기 ...
    next_token = torch.multinomial(softmax(logits), 1) if do_sample else argmax(logits)
    input_ids = torch.cat([input_ids, next_token], -1)
    past_key_values = outputs.past_key_values if use_cache else None

KV 缓存。第一步一次性放入 P 个提示令牌(prefill)。有缓存的话,之后每一步只放入一个令牌,所以生成 N 个的过程中,进入 forward 的令牌是 P + (N − 1) 个。没有缓存的话,每一步都要重新放入 P、P+1、……,总共是 P·N + N(N−1)/2 个。提示 183 个令牌、生成 64 个令牌时,是 246 对 13,728——相差 56 倍。多亏有因果掩码,前面令牌的 K、V 不会因为后面令牌的出现而改变,所以取出来用,结果也相同。

缓存的大小。缓存是每层的 K、V,形状是(批次,长度,KV 头,head_dim)。MiniMind 把用 repeat_kv 复制 之前 的 K、V 放入缓存,所以 GQA 减少了多少,它就原样减少多少。100 个令牌就是 2 × 4 层 × 100 × 2 个头 × 32 × 4 字节 = 204,800 字节。在服务中,这个数字决定并发用户数。

温度。把 logits 除以温度之后再做 softmax。小于 1 时,分布变尖,集中到最可信的令牌上;大于 1 时,分布变平,罕见的令牌也会被选中。温度接近 0,就和贪心解码相同。

top-p(nucleus)。按概率从大到小排序,只保留到累积概率超过 p 的位置,其余的用 −∞ 抹掉。MiniMind 把掩码错开一位,保留 直到第一次超过 p 的那个令牌,并且最上面的令牌总是保留。分布很尖时只留几个,分布平时留很多——这与固定保留个数的 top-k 不同。MiniMind 的 generate 默认值是温度 0.85、top_p 0.85、top_k 50。

在现场相遇的样子

“打开缓存之后答案变了”这类缺陷报告,多半不是缓存的问题,而是采样一侧的问题——没有固定随机数,或者只有一边设了 do_sample。先用贪心解码比较两种方式,看是否一个令牌都不差,就能把原因范围缩小一半。反过来,总结长文档时,到第一个令牌为止的时间(TTFT)很长,缓存是减不掉的——prefill 本来就得把提示整个算一遍。缓存减少的是之后令牌之间的时间。

本课程与 MiniMind 原版的不同之处

MiniMind 的 eval_llm.py 和网页演示用温度 0.85、top_p 0.95 来取令牌,还可以设置重复惩罚(repetition_penalty)。这门课程为了让实验干净,把 top_k 关掉(top_k=0),每次只改一个因素。另外,MiniMind 的 generate 会按批次分别记住已经结束的行(finished),全部结束时才停下,对已结束的行则一直填充 eos。这就是一次取 30 个时,短回答后面接着一串 eos 的原因。真实的服务引擎(vLLM 等)还在此基础上增加了把每个请求长度不同的缓存捆在一个批次里的装置——原理相同,只是管理变得复杂。

下一项实验要做什么

用基准预训练模型,看有缓存和无缓存时是否产生相同的令牌,数出进入 forward 的令牌数并与公式对上,再测量时间。让缓存的实际字节数与公式对上,用 SFT 模型改变温度,每次取 30 次,数出不同回答的个数,并数出 top-p 保留的候选数。