自回归推理与 KV-Cache 显存优化:打造毫秒级打字机

从零手写工业级大模型自回归推理引擎:剖析失忆金鱼 O(N^2) 重复计算灾难、KV-Cache 便签记忆法与 6.29MB 极致轻量显存账本、深入拆解解码采样四大黄金参数(重复惩罚、温度系数、Top-K 截断与 Top-P 核采样向右平移一位魔法),并结合 TextGenerator 实现丝滑终端流式打字机。

本文目录27 个章节

本章导读: 在前面的章节中,我们已经训练好了模型的“大脑”。但此时的模型就像一个刚刚读完万卷书的学生,还坐在考场上没有动笔。 怎么让它开口说话?怎么让它像 ChatGPT 一样,文字如泉涌般、像打字机一样一个字一个字向外跳动? 很多初学者写出自回归推理代码后,往往会发现两个令人沮丧的现象:第一,句子越长,生成速度越慢,最后卡得像幻灯片;第二,模型说起话来颠三倒四,或者像个复读机一样反复念叨一句话。 本章将像老师坐在你身边一样,不甩冷冰冰的代码,而是通过逐行拆解、具体数字推演和生活比喻,带你从零写出一个工业级水准的自回归推理引擎。

5.1 什么是自回归生成?从“成语接龙”讲起

5.1.1 一个小学生都能听懂的自回归过程

所谓“自回归 (Autoregressive)”,名字听起来非常高深,但它的生活原型就是我们小时候都玩过的游戏——文字接龙

想象你给模型出了一道上文(这在 AI 里叫 Prompt,提示词):

“今天天气真好,我们一起去”

模型的大脑中并没有一个现成的句子库,它的写作过程是极其纯粹的逐字蹦豆子

文本
第 1 轮思考:
  输入上文: "今天天气真好,我们一起去" (共 11 个字)
  模型计算: 算出了词表里 4,096 个词各自出现的概率
  挑选结果: 概率最高的是 "公"
  输出文字: "公"

第 2 轮思考:
  把刚刚写出来的 "公" 拼到原来的上文屁股后面!
  新的上文: "今天天气真好,我们一起去公" (变成 12 个字)
  模型计算: 重新算一遍 4,096 个词的概率
  挑选结果: 这一次最合适的是 "园"
  输出文字: "园"

第 3 轮思考:
  新的上文: "今天天气真好,我们一起去公园" (变成 13 个字)
  模型计算: 重新算概率...
  挑选结果: 最合适的是 "散"
  输出文字: "散"

看到规律了吗?自己刚刚写出来的字(输出),立刻掉头变成下一轮思考的原料(输入),这就叫“自”(自己)“回归”(回到输入)!

5.2 朴素推理的致命陷阱:失忆金鱼的 O(N2)O(N^2) 灾难

如果你按照上面这个最直观的逻辑去写代码,就会掉进计算机科学中最著名的性能深渊。

5.2.1 患有强迫症的金鱼打字员

想象一个很有礼貌但记忆力极差的打字员:

  • 他每敲出一个新字,他的老板要求他必须把前面已经写好的所有文字从第一个字开始,大声重新朗读一遍
文本
===================================================================================
                       朴素自回归推理的重复计算噩梦
===================================================================================
生成第 1 个字 ("公"):
  打字员朗读上文: "今天天气真好,我们一起去"                  ──> 算了 11 个字的特征

生成第 2 个字 ("园"):
  打字员从头重读: "今天天气真好,我们一起去公"                ──> 算了 12 个字的特征

生成第 3 个字 ("散"):
  打字员再次从头重读: "今天天气真好,我们一起去公园"          ──> 算了 13 个字的特征
  ...
生成第 100 个字:
  打字员又从头重读: [ 前面全部 110 个字全部重读一遍! ]       ──> 算了 111 个字的特征
===================================================================================

5.2.2 算一笔恐怖的账

如果我们要生成一篇 500 字的小短文: 总计算量=11+12+13++511(11+511)×5002=130,500 次单词计算\text{总计算量} = 11 + 12 + 13 + \dots + 511 \approx \frac{(11 + 511) \times 500}{2} = \mathbf{130,500\text{ 次单词计算}}! 而如果打字员拥有正常的记忆力,每一步只算新来的那 1 个字,总共只需要算 500500! 朴素的做法竟然平白无故多算了整整 260 倍

推理模式类型每轮迭代送入模型的输入时间复杂度生成 500 字计算量实际使用体验
朴素无缓存推理 (Naive)全部历史上下文(长度为 TT 的完整长句)O(N2)O(N^2)(二次方膨胀)130,500130,500 次前向特征计算随生成文本变长,字与字之间停顿越来越长,最后严重卡顿。
KV-Cache 极速推理 (工业标配)仅仅送入最新生成的单字(长度恒为 1)O(N)O(N)(线性恒定)仅需 500500 次前向特征计算(提速 260 倍!无论生成多长的文章,打字机始终如丝般顺滑、均匀吐字。

这就是为什么很多新人自己写推理时,生成前 5 个字挺快,越往后越慢,最后 GPU 风扇狂转、显卡几乎瘫痪的原因。

5.3 破局之刃:手把手教你理解 KV-Cache 的“便利贴记忆法”

怎么解决这个重复计算的灾难?核心答案就是:KV-Cache(键值缓存)

5.3.1 为什么历史上的词不需要重算?

回想我们在第 02 章讲过的自注意力机制公式: Attention(Q,K,V)=softmax(QKTdk)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{Q K^T}{\sqrt{d_k}}\right) V

请静下心来仔细思考这三个字母分别代表什么:

  • QQ (Query,提问):当前文字“想问什么”。
  • KK (Key,标签):每个历史文字“包含什么主题”。
  • VV (Value,内容):每个历史文字“实际承载的信息”。

当我们在写第 100 个字时:

  1. 谁是新的? 只有第 100 个字是刚刚诞生的,所以唯独需要计算第 100 个字的提问向量 q100q_{100}
  2. 谁是旧的? 前面 99 个字早就写在纸上了!第 1 个字的 k1,v1k_1, v_1,第 2 个字的 k2,v2k_2, v_2……它们早在几十秒之前就已经计算过了!
  3. 因果关系的铁律:因为大模型是自回归因果模型,未来的词绝不能倒流影响过去。也就是说,第 1 个字身后的 k1k_1v1v_1,在整个宇宙的时间线上永远不会再发生任何改变!

5.3.2 便利贴工作法图解

既然旧词的 KKVV 永远不会变,我们何必每次傻乎乎地从头重算呢? 我们只需准备一个小文件夹(这就是 Cache 缓存),把每一层算出来的 KKVV 像便利贴一样贴在里面!

流程图

奇迹发生了:无论文章写到第 10 个字还是第 1000 个字,每一次预测都只消耗 1 个字的前向计算量!打字机的响应速度始终均匀、稳定、恒定如初!

5.3.3 算一笔账:0.04B 架构的 KV-Cache 到底占多少显存?

很多人听说大模型推理吃显存,那么对于我们精心设计的 0.04B 黄金架构,维护这套 KV-Cache 便签条需要多少字节? 我们来精确计算单次会话(Batch Size = 1)在满血 1,024 上下文长度下的显存占用:

  • 网络总层数 nlayers=12n_{\text{layers}} = 12
  • 我们采用 GQA 分组查询注意力,每层仅有 nkv=2n_{kv} = 2 个键值头(相比传统 MHA 的 8 个头,显存暴降 75%!);
  • 每个头的特征维度 dhead=64d_{\text{head}} = 64
  • 采用半精度浮点数(fp16,每个数字占 2 字节);
  • 同时需要缓存 KKVV 两个张量。

每一个 Token 所产生的 KV 缓存大小:

Memorytoken=2×12 (层)×2 (KV头)×64 (维度)×2 (Bytes)=6,144 字节6.14 KB\text{Memory}_{\text{token}} = 2 \times 12\text{ (层)} \times 2\text{ (KV头)} \times 64\text{ (维度)} \times 2\text{ (Bytes)} = \mathbf{6,144\text{ 字节}} \approx \mathbf{6.14\text{ KB}}

跑满极限 1,024 完整上下文的显存开销:

Memorytotal=1,024×6,144 字节6.29 MB\text{Memory}_{\text{total}} = 1,024 \times 6,144\text{ 字节} \approx \mathbf{6.29\text{ MB}}!

[!TIP] 结论极其震撼:即使生成写满整整 1,024 个字的大长篇,整个 0.04B 模型的 KV-Cache 仅仅需要 6.29 MB 的极微弱显存!相比全注意力(MHA)的 25.2 MB 降低了 75%,即使在手机、树莓派等极端嵌入式边缘设备上也能毫无压力地以毫秒级极速流畅推理!

5.4 采样策略逐行精讲:如何让模型不再说胡话?

模型通过神经网络计算完毕后,吐出的不是汉字,而是一个长度为 4,096 的一维数字向量,在数学上称为 未归一化得分 (Logits)。 如何把这 4,096 个打分转化成最终说出口的那一个字?这完全取决于采样策略 (Sampling Strategy)

很多教程直接贴出一大段采样代码,新人往往看得一头雾水。在进入逐行推演前,先来看这四大金牌参数的职责总览:

采样超参数典型推荐值调控物理含义调高/调低实际体验对比
重复惩罚 (repetition_penalty)1.1 ~ 1.2对近期已生成的词强行打折降分设为 1.0:无惩罚,容易陷入死循环复读;<br>设为 1.2:有效破除车轱辘话,词汇多样性提升。
生成温度 (temperature)0.7 ~ 0.8缩放原始打分差距,控制概率平坦度低(0.1~0.5):严谨确定、聚焦高分词,适合对仗格律;<br>高(1.0~1.5):思维发散、富于创意,但过高会胡言乱语。
候选截断 (top_k)40 ~ 50暴力截断法:仅保留得分最高的前 K 个词设为 0:不截断;<br>设为 50:一刀切除后 4000+ 个完全不相干的低分怪词。
核采样 (top_p)0.8 ~ 0.9动态质量法:保留累积概率达 P 的动态核心池随上下文自适应:上文确定时(如成语)候选池极小;上文开阔时候选池自动扩大,比 Top-K 更聪明。

下面我们用一个具体的小例子,逐行带你推演采样函数的每一个微小细节

假设场景

假设模型的词表里现在只有 4 个备选词,它们的原始打分分别为:

  • "北京": 4.0 分
  • "上海": 3.0 分
  • "广州": 2.0 分
  • "火星": 0.1 分 (不靠谱的荒谬选项)

5.4.1 第一步:重复惩罚 (Repetition Penalty) —— 禁止复读机

如果模型刚刚上一句话已经说了 "北京",我们希望它不要翻来覆去念叨。

核心逻辑代码(在 src/generate.py 中,我们还特意对标点符号与特殊符进行豁免,防止标点因惩罚产生交替震荡):

PYTHON
if repetition_penalty != 1.0 and generated_ids:
    exempt_ids = {0, 1, 2, 3, 14, 262, 263, 273, 274, 275, 534, 1227, 1916}
    for tid in set(generated_ids):
        if tid in exempt_ids:
            continue
        if logits[tid] < 0:
            logits[tid] *= repetition_penalty
        else:
            logits[tid] /= repetition_penalty

小白推演

  • 假设惩罚系数设为 1.2(表示只要说过就扣分);
  • "北京" 的得分是正数 4.0,执行 4.0 / 1.2 = 3.33!它的得分被生生削弱了;
  • 如果某个词本来就是负分(比如 -2.0),乘上 1.2 变成 -2.4,同样也变得更低。
  • 这样一来,模型选到重复词的倾向被大大抑制,有效根治“车轱辘话”。

5.4.2 第二步:温度调节 (Temperature) —— 注入灵性还是保持严谨?

核心逻辑代码

PYTHON
logits = logits / temperature

就这么短短一行代码,为什么被称为“温度”?我们拿数字做个实验:

情况 A:温度很高 (T=2.0T = 2.0,让模型放飞自我)

  • "北京": 4.0/2.0=2.04.0 / 2.0 = \mathbf{2.0}
  • "上海": 3.0/2.0=1.53.0 / 2.0 = \mathbf{1.5}
  • "广州": 2.0/2.0=1.02.0 / 2.0 = \mathbf{1.0}
  • "火星": 0.1/2.0=0.050.1 / 2.0 = \mathbf{0.05}
  • 效果:原本差距悬殊的 4 分和 2 分,现在变成了 2 分和 1 分,差距被缩小拉平了!所有的词概率变得更加均匀,模型开始天马行空,富有创造力,但也容易胡思乱想。

情况 B:温度很低 (T=0.5T = 0.5,让模型严谨冷静)

  • "北京": 4.0/0.5=8.04.0 / 0.5 = \mathbf{8.0}
  • "上海": 3.0/0.5=6.03.0 / 0.5 = \mathbf{6.0}
  • "广州": 2.0/0.5=4.02.0 / 0.5 = \mathbf{4.0}
  • 效果:原本的微小差距被极度放大!高分选手骑绝尘,低分选手彻底丧失机会。模型会极度自信、严谨、确定。

5.4.3 第三步:Top-P (Nucleus) 核采样 —— 为什么需要移位操作?

这是让无数新手最抓狂的一段经典代码。很多人在看开源库时,看到下面这几行代码,感觉宛如天书:

PYTHON
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)

# 下面这两行到底在干嘛?!
sorted_indices_to_remove = cumulative_probs > top_p
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
sorted_indices_to_remove[..., 0] = 0

别慌!老师现在用最通俗的推演,把这几行代码解剖给你看:

场景模拟推演

假设 4 个词经过 Softmax 后的概率如下:

  1. "北京": 概率 60%
  2. "上海": 概率 25%
  3. "广州": 概率 10%
  4. "火星": 概率 5%

我们希望设定 top_p = 0.8(80%),也就是只在最有希望的前 80% 优质词汇里挑选,把剩下的冷门烂词(如“火星”)剔除

1. 计算累加概率 (cumulative_probs):

  • "北京": 60%60\%
  • "上海": 60%+25%=85%60\% + 25\% = \mathbf{85\%} (此时累加已经跨过了 80% 的门槛!)
  • "广州": 85%+10%=95%85\% + 10\% = 95\%
  • "火星": 95%+5%=100%95\% + 5\% = 100\%

2. 如果直接使用 cumulative_probs > 0.8 会发生什么?

  • "北京" (60%): 没超过 0.8 \to 保留
  • "上海" (85%): 超过了 0.8 \to 竟然被标记为剔除!
  • "广州" (95%): 剔除
  • "火星" (100%): 剔除
  • 致命问题暴露了"上海" 是让总概率刚刚达到 80% 必不可少的功臣!如果把它剔除掉,候选池里就只剩孤零零的一个 "北京"(才 60%),根本凑不够我们要求的 80% 候选面!

3. 神奇的“向右平移一位”解法!

为了把刚好跨过阈值的这个“临界词”("上海")安全保下来,工程师想出了一个极其精妙的数学技巧——把所有标记向右平移一位

TEXT
原始累加是否超过 0.8:  [ False,   True,   True,   True ]
                                   ↘       ↘       ↘
向右平移一位后结果:    [   ?,   False,   True,   True ]
强制将第一名设为 False: [ False,  False,   True,   True ]
                         (北京)   (上海)   (广州)  (火星)
                          保留     保留     剔除    剔除!

看!"上海" 被奇迹般地挽救了回来! 候选池里留下了 "北京" (60%) + "上海" (25%) = 85%,完美覆盖了 80% 的概率空间,而垃圾长尾词 "广州""火星" 被干净利落地切除了! 现在你再回头看那两行代码:

PYTHON
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone() # 向右平移
sorted_indices_to_remove[..., 0] = 0 # 无论如何,绝对不能把第一名剔除掉

是不是顿时恍然大悟?这就是顶级架构设计的数学美感!

5.5 完整的自回归采样函数拆解

有了上面的全部铺垫,现在我们把完整的采样逻辑组装起来。每一行都有详细注释:

PYTHON
import torch
import torch.nn.functional as F
from typing import List, Optional

def sample_next_token(
    logits: torch.Tensor,
    temperature: float = 0.8,
    top_k: int = 50,
    top_p: float = 0.9,
    repetition_penalty: float = 1.1,
    generated_ids: Optional[List[int]] = None
) -> int:
    """
    自回归核心采样函数:将模型打分转化为最终选出的单字 ID
    """
    # 保证操作不影响原张量,且展平为一维 [Vocab_Size]
    logits = logits.clone().squeeze()

    # 1. 重复惩罚:削减历史出现过文字的分数,拒绝复读机
    if repetition_penalty != 1.0 and generated_ids:
        for tid in set(generated_ids):
            if logits[tid] < 0:
                logits[tid] *= repetition_penalty
            else:
                logits[tid] /= repetition_penalty

    # 2. 极端贪心搜索保护:如果温度无限接近 0,直接取第一名
    if temperature <= 1e-4:
        return torch.argmax(logits, dim=-1).item()

    # 3. 温度缩放:平滑或陡峭化概率曲线
    logits = logits / temperature

    # 4. Top-K 截断:只保留前 K 个最高分的强力候选
    if top_k > 0:
        top_k = min(top_k, logits.size(-1))
        # 找出第 K 名的分数,凡是比它小的全部赋为负无穷大 (-inf)
        k_th_val = torch.topk(logits, top_k)[0][..., -1, None]
        indices_to_remove = logits < k_th_val
        logits[indices_to_remove] = -float("Inf")

    # 5. Top-P 核采样:动态累加概率,切除长尾荒谬候选
    if top_p < 1.0:
        # 从大到小排序
        sorted_logits, sorted_indices = torch.sort(logits, descending=True)
        # 计算累加概率和
        cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)

        # 找到超出阈值的截断点
        sorted_indices_to_remove = cumulative_probs > top_p
        # 巧妙右移,保留临界词
        sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
        sorted_indices_to_remove[..., 0] = 0

        # 将被剔除的索引还原映射回原 logits,并置为负无穷大
        indices_to_remove = sorted_indices[sorted_indices_to_remove]
        logits[indices_to_remove] = -float("Inf")

    # 6. 转为概率分布,并进行轮盘赌抽奖 (Multinomial 依概率采样)
    probs = F.softmax(logits, dim=-1)
    next_token_id = torch.multinomial(probs, num_samples=1).item()
    return next_token_id

5.6 动手实现流式生成器:打造你的专属打字机

现在我们要写一个类:TextGenerator。它把分词器、模型、KV-Cache 和采样函数全部串联在一起,实现真正的终端流式输出。 源码位于 src/generate.py

逐段代码讲解与实现:

PYTHON
class TextGenerator:
    def __init__(self, model, tokenizer, device: str = "cpu"):
        self.model = model.to(device).eval() # 切换到评估模式,关闭 Dropout
        self.tokenizer = tokenizer
        self.device = device

    @torch.no_grad() # 推理阶段绝不计算梯度,节约海量显存与计算时间
    def generate(
        self,
        prompt: str,
        max_new_tokens: int = 128,
        temperature: float = 0.8,
        top_k: int = 50,
        top_p: float = 0.9,
        repetition_penalty: float = 1.1,
        stream: bool = True
    ) -> str:
        # 第一步:把人类输入的提示词编码成数字序列
        prompt_ids = self.tokenizer.encode(prompt, add_bos=True)
        generated = list(prompt_ids)

        if stream:
            print(prompt, end="", flush=True)

        past_key_values = None
        input_ids = torch.tensor([prompt_ids], dtype=torch.long, device=self.device)

        # 第二步【首字填充 Prefill 阶段】:
        # 一次性将完整的 Prompt 送入模型,产生初始的 KV-Cache 便签条,并得到第 1 个新字
        logits, _, past_key_values = self.model(input_ids, past_key_values=None, start_pos=0)
        next_token_logits = logits[0, -1, :] # 提取 Prompt 最后一个位置的预测得分

        # 采样出第 1 个生成的词
        next_id = sample_next_token(
            next_token_logits,
            temperature=temperature,
            top_k=top_k,
            top_p=top_p,
            repetition_penalty=repetition_penalty,
            generated_ids=generated
        )

        generated.append(next_id)
        if stream:
            # 实时将这个词翻译回字符,并打在终端屏幕上!
            print(self.tokenizer.decode([next_id], skip_special_tokens=False), end="", flush=True)

        start_pos = len(prompt_ids)

        # 第三步【增量自回归 Decode 阶段】:
        # 依靠 KV-Cache,每次只喂入最新生成的 1 个字!
        for _ in range(max_new_tokens - 1):
            # 如果生成了句子结束符 </s>,说明文章写完了,优雅退出
            if next_id == self.tokenizer.eos_token_id:
                break
            # 如果长度超出了模型最大上限,停止生成
            if start_pos >= self.model.config.max_seq_len:
                break

            # 核心提速秘诀:单字输入!形状仅仅为 [1, 1]
            step_input = torch.tensor([[next_id]], dtype=torch.long, device=self.device)
            logits, _, past_key_values = self.model(
                step_input,
                past_key_values=past_key_values, # 传入并持续追加便利贴
                start_pos=start_pos
            )
            next_token_logits = logits[0, -1, :]

            # 采样下一个字
            next_id = sample_next_token(
                next_token_logits,
                temperature=temperature,
                top_k=top_k,
                top_p=top_p,
                repetition_penalty=repetition_penalty,
                generated_ids=generated
            )

            generated.append(next_id)
            start_pos += 1

            if stream:
                print(self.tokenizer.decode([next_id], skip_special_tokens=False), end="", flush=True)

        if stream:
            print() # 换行收尾

        # 返回全部新生成的文字内容
        return self.tokenizer.decode(generated[len(prompt_ids):])

5.7 本节核心收获与课后思考

通过本章的深入学习,你已经彻底搞懂了支撑 ChatGPT 交互体验的最底层逻辑:

  1. 自回归的本质:是上一步生成的文字掉头成为下一步输入的接龙循环;
  2. KV-Cache 的魔力:利用因果注意力历史不变的特性,将旧词的 K,VK, V 永久缓存,一举消灭 O(N2)O(N^2) 计算量灾难;
  3. 采样的艺术
    • 想要做严谨客观的数学题,把温度调到接近 0;
    • 想要写诗、讲童话故事,把温度设为 0.8,搭配 Top-P 0.9 过滤杂音;
    • 想要杜绝车轱辘话,开启重复惩罚系数 1.1。

💡 课后动手小实验

试着修改 scripts/run_generate.py 中的参数:

  1. temperature 设为 2.0,看看模型会说出怎样荒诞不经的怪话?
  2. repetition_penalty 设为 1.0(完全不惩罚),看看模型是不是更容易陷入复读机模式?

在下一章中,我们将进行整套系统的全流程实战与总装测试——第 06 章|0.04B 模型全流程端到端实战贯通与评估

REFERENCES

参考链接

  1. 01The Curious Case of Neural Text Degeneration - Nucleus Sampling (Holtzman et al.)
  2. 02GQA: Training Generalized Multi-Query Transformer Models (Ainslie et al.)
  3. 03vLLM High-Throughput PagedAttention Serving

所属系列

从零开始手搓大模型

下一步

继续浏览相关主题

沿着同一主题继续阅读。

查看最新资讯