系列文章:AI 大模型知识体系 | 第三周・第一篇


引言:模型训完了,然后呢?

恭喜你!经过前两周学习,我们吃透了 Transformer 架构、预训练、SFT、RLHF、分布式训练…… 模型训练完成,权重文件也保存完毕,但这只是第一步。

真正的性能挑战,从推理阶段才开始。

训练大模型好比打造一台专业赛车:投入大量算力、资金、时间完成制作,但赛车不能只放在车库封存,必须开上赛道跑起来。推理(Inference) 就是模型 “上路运行” 的过程:接收用户输入,输出通顺回答,同时满足低延迟、高吞吐、低显存消耗三大要求。

很多新手误以为推理只是简单调用model.generate(),实际差距天差地别:

  1. 同一款 7B 大模型,无优化每秒仅生成 5 个 token,搭配 KV Cache 等优化后可达 100+ token;

  2. 70B 超大模型,不开启 KV Cache 时,生成第 1000 个 token 耗时是第 1 个的 100 倍;

  3. 推理阶段显存瓶颈、优化逻辑,和模型训练完全割裂。

本文是大模型推理入门核心教程,重点拆解 KV Cache 原理,它是所有推理加速方案的底层基石。


一、推理和训练到底有什么区别?

不少人混淆训练与推理,认为二者只是单纯跑模型,二者差异等同于学车全职开出租,核心对比表如下:

对比维度

训练

推理

核心目标

学习数据分布,更新模型参数

复用已训练参数,生成文本结果

生活化类比

学生刷题、课堂学习

老司机接单载客

计算链路

前向传播 + 反向传播

仅前向传播,无梯度计算

批量大小

越大越好,最大化 GPU 算力利用率

批量通常很小,线上多为单请求 batch=1

性能瓶颈

计算密集(矩阵乘法耗时最高)

显存密集(KV Cache 占用绝大部分显存)

数值精度

高精度 FP32/BF16,保证梯度稳定

低精度 INT8/INT4,不影响阅读语义即可

配套组件

优化器 AdamW、梯度缓存、学习率调度

无优化器、无梯度缓存

成本特征

一次性高额算力开销

单次调用成本极低,但请求总量巨大

核心结论:训练吃算力,推理吃显存。训练时 GPU 算力占满,显存还有富余;推理时常出现算力闲置、显存直接爆显存的情况,这也决定两者优化思路完全相反。


二、大模型推理两大阶段:Prefill 预填充 & Decode 解码

文本生成流程分为两个独立阶段,用人写回答类比:

Prefill(预填充)= 一次性读完用户全部提问 Decode(解码)= 逐字输出回答内容

2.1 Prefill:并行一次性处理完整 Prompt

用户输入长文本 Prompt 后,模型第一步统一处理全部输入 token,完整理解上下文信息,这个阶段叫 Prefill。

  • 特征:全部 token 并行计算,单次执行,计算量大;

  • 场景举例:用户输入 2000 字长文要求总结,模型一次性计算全文所有 token 的 K、V 向量。

逻辑示意:


用户输入:请帮我写一首七言绝句春日诗 ↓ 【Prefill阶段】 一次性并行计算全部输入token的K、V向量,完整缓存上下文 输入序列:[请,帮,我,写,一,首,...] 并行计算全部token K、V,一次性完成 特点:计算密集,仅执行一次

2.2 Decode:逐 token 生成输出文本

Prefill 完成后,模型循环生成回答,每轮仅输出 1 个 token,每一步都需要依赖输入 + 已生成全部文本做注意力计算,该阶段为 Decode。

逻辑示意:


【Decode阶段】循环执行,每次仅生成1个token Step1:输入文本 + [春] → 预测下一字「风」 Step2:输入文本 + [春,风] → 预测下一字「拂」 Step3:输入文本 + [春,风,拂] → 预测下一字「面」 循环直至触发停止符 痛点:每一步都需要重新计算历史所有token的K、V向量,大量重复计算,速度极慢

原生 Decode 致命缺陷

无缓存机制时,生成 N 个 token 的注意力计算总次数为 1+2+3+...+N,序列越长开销爆炸。输出 1000token 需要近 50 万次重复注意力运算,完全无法落地线上服务。 KV Cache 就是为消除重复计算而生。


三、KV Cache:推理加速专属 “读书笔记”

3.1 KV Cache 核心原理

回顾 Self-Attention 基础:每个 token 会映射出三组向量 Q(查询)、K(键)、V(值),注意力分数由 Q 与全部 K 匹配,加权求和 V 得到输出。

KV Cache 核心逻辑极简:已经计算完成的历史 token K、V 向量,存入显存缓存,后续步骤直接读取,不再重复计算

生活化类比:阅读长篇小说,每次续写都要回顾前文。无缓存 = 每次从头重读全书;KV Cache = 提前记录前文关键信息,新增内容只需要读取笔记,不用重复翻书。

两种模式对比:


无KV Cache(重复计算) 生成token1:计算K1、V1 生成token2:重新计算K1,V1 + 新K2,V2 生成token3:重新计算K1,V1、K2,V2 + 新K3,V3 生成第N个token:重复计算全部N组KV,计算量平方增长 开启KV Cache(复用缓存) 生成token1:计算K1,V1 → 存入缓存 生成token2:读取缓存K1V1,仅计算K2V2,拼接更新缓存 生成token3:读取历史全部缓存,仅计算新增KV 生成第N个token:只算最新1组KV,历史全部复用

3.2 理论与实际加速效果

  1. 计算复杂度变化:无缓存O(N²)平方级开销 → KV Cache O(N)线性开销,长文本提升极其明显;

  2. 代码极简模拟(PyTorch)


import torch import time # 无缓存注意力:每轮重算全部KV def attention_no_cache(Q, K, V): scores = torch.matmul(Q, K.transpose(-2, -1)) / (K.size(-1) ** 0.5) weights = torch.softmax(scores, dim=-1) return torch.matmul(weights, V) # 带KV Cache注意力:复用历史KV,仅拼接新增 def attention_with_cache(new_Q, cached_K, cached_V, new_K, new_V): K = torch.cat([cached_K, new_K], dim=1) V = torch.cat([cached_V, new_V], dim=1) scores = torch.matmul(new_Q, K.transpose(-2, -1)) / (K.size(-1) ** 0.5) weights = torch.softmax(scores, dim=-1) out = torch.matmul(weights, V) return out, K, V # 测试参数 seq_len = 1024 d_k = 64 num_heads = 32 # 无缓存耗时测试 start = time.time() K_all = torch.randn(1, seq_len, num_heads, d_k) V_all = torch.randn(1, seq_len, num_heads, d_k) for i in range(seq_len): Q = torch.randn(1, 1, num_heads, d_k) attention_no_cache(Q, K_all[:, :i+1], V_all[:, :i+1]) time_no_cache = time.time() - start print(f"无KV Cache耗时: {time_no_cache:.3f}s") print(f"KV Cache理论加速比:约{seq_len}倍,序列越长提升越大")

  1. 真实落地效果:LLaMA-7B 生成长度 512 文本,KV Cache 可带来10~20 倍生成速度提升,是线上服务必备基础优化,不开启完全无法商用。


四、KV Cache 显存开销:速度的代价是显存占用

KV Cache 解决重复计算,但会持续占用显存存储每一层、每个 token 的 K、V 向量,长序列、多并发场景显存占用压力极大。

4.1 KV Cache 显存计算公式

参数说明:

  • 2:K 向量、V 向量两份存储;

  • dtype 字节:FP16/BF16=2Byte,INT8=1Byte,INT4=0.5Byte;

  • 序列长度 = 输入 Prompt 长度 + 已生成输出长度。

4.2 主流模型显存占用实测(batch=1,FP16)

模型

层数

KV 头数

单头维度

序列 2048

序列 4096

序列 8192

LLaMA-7B

32

32

128

1.0GB

2.0GB

4.0GB

LLaMA-13B

40

40

128

1.6GB

3.2GB

6.4GB

LLaMA-70B

80

64

128

6.4GB

12.8GB

25.6GB

计算示例:LLaMA-7B,seq_len=2048,FP16,单请求

4.3 多并发显存压力示例


def calc_kv_cache_memory(num_layers, num_kv_heads, head_dim, seq_len, batch_size=1, dtype_bytes=2): total_bytes = 2 * num_layers * seq_len * num_kv_heads * head_dim * dtype_bytes * batch_size return total_bytes / (1024 ** 3) # LLaMA-7B 单用户、多用户对比 print(f"LLaMA-7B seq=2048 单用户:{calc_kv_cache_memory(32,32,128,2048):.2f}GB") print(f"LLaMA-7B seq=2048 10并发:{calc_kv_cache_memory(32,32,128,2048,batch_size=10):.2f}GB")

输出结果:


LLaMA-7B seq=2048 单用户:1.00GB LLaMA-7B seq=2048 10并发:10.00GB

10 个用户同时对话,仅缓存就要占用 10GB 显存,这也是大模型推理显存紧张的核心根源。


五、KV Cache 显存瘦身方案:MHA / MQA / GQA

显存占用来自 KV 多头独立存储,优化核心思路:减少独立 KV 头数量,多头共享 KV 向量

5.1 MHA 标准多头注意力(基准方案)

原始 GPT、初代 LLaMA 采用,每个 Q 头配套独立 K、V 头:

  • Q:32 独立查询头

  • K:32 独立键头

  • V:32 独立值头

  • 缓存体积:1 倍基准,显存开销最大,精度最优

5.2 MQA 多查询注意力

极致压缩方案,全部 Q 头共享唯一一组 K、V:

  • Q:32 独立查询头

  • K/V:全局仅 1 组共享向量

  • 缓存体积:仅原生 MHA 的 1/32,显存大幅降低

  • 缺点:KV 多样性不足,长文本、复杂逻辑场景精度轻微下降;代表模型:Falcon、PaLM

5.3 GQA 分组查询注意力(工业主流)

MHA 与 MQA 折中方案,将 Q 头划分为多组,每组共享一套 KV 向量,平衡精度与显存开销,LLaMA2/3、Mistral 全系标配。 举例:32 个 Q 头分为 4 组,每组 8 个 Q 头共用一套 KV,仅需 4 组 KV 存储,缓存体积降至原生 1/8。

三者横向对比

方案

Q 头

KV 头

缓存体积

文本精度

推理速度

代表模型

MHA

32

32

1.0× 基准

最高

最慢

GPT3、LLaMA1

GQA

32

8

0.25×

接近 MHA

LLaMA2/3、Mistral

MQA

32

1

0.03×

轻微下降

最快

Falcon、PaLM

显存差距量化示例(LLaMA-70B seq=4096)


def mha_kv(layers, heads, head_dim, seq): return 2 * layers * seq * heads * head_dim * 2 / (1024**3) def gqa_kv(layers, q_heads, kv_heads, head_dim, seq): return 2 * layers * seq * kv_heads * head_dim * 2 / (1024**3) def mqa_kv(layers, head_dim, seq): return 2 * layers * seq * 1 * head_dim * 2 / (1024**3) layers, q_heads, head_dim, seq = 80, 64, 128, 4096 mha = mha_kv(layers, q_heads, head_dim, seq) gqa = gqa_kv(layers, q_heads, 8, head_dim, seq) mqa = mqa_kv(layers, head_dim, seq) print(f"MHA缓存:{mha:.2f}GB") print(f"GQA缓存:{gqa:.2f}GB,节省{(1-gqa/mha)*100:.0f}%显存") print(f"MQA缓存:{mqa:.2f}GB,节省{(1-mqa/mha)*100:.0f}%显存")

输出:


MHA缓存:12.80GB GQA缓存:1.60GB,节省88%显存 MQA缓存:0.20GB,节省98%显存


六、三大核心推理性能指标

线上服务评判模型推理效果,依靠三个标准指标:

6.1 TTFT 首 token 延迟

从发送请求到返回第一个文字的耗时,主要衡量 Prefill 阶段性能。 类比:餐厅下单到第一道菜上桌的等待时间,TTFT 过高用户会感知模型卡顿。 影响因素:Prompt 长度、模型参数量、Prefill 并行优化程度。

6.2 TPS 每秒生成 token 数

Decode 阶段核心指标,代表文字输出速度,数值越高,打字流畅度越好。 影响因素:KV Cache、GQA、批处理、显存带宽。

6.3 端到端总延迟

完整生成全部回答的总耗时,计算公式:

场景对比代码示例


def compute_metrics(ttft_ms, tps, out_tokens): decode_ms = out_tokens / tps * 1000 total_ms = ttft_ms + decode_ms print(f"首字延迟TTFT:{ttft_ms:.0f}ms") print(f"生成速度TPS:{tps:.1f} token/s") print(f"解码耗时:{decode_ms:.0f}ms") print(f"端到端总延迟:{total_ms:.0f}ms\n") # 无优化原生推理 print("【无KV Cache原生推理】") compute_metrics(800, 8, 256) # 基础KV Cache+GQA print("【KV Cache+GQA优化】") compute_metrics(300, 60, 256) # vLLM分页缓存全套优化 print("【vLLM全套优化】") compute_metrics(150, 120, 256)

输出效果差距极大,无优化需要 32 秒生成回答,极致优化仅需 2 秒左右。


七、推理阶段两大性能瓶颈

推理分 Prefill、Decode 两个阶段,瓶颈类型完全不同,优化方向区分开:

7.1 Prefill:计算密集型 Compute-Bound

一次性并行计算大量 token,矩阵乘算力占用满,显存带宽富余。 优化方向:算子融合、FlashAttention、Tensor Core 加速。

7.2 Decode:显存带宽密集型 Memory-Bound

每轮仅生成 1 个 token,算力闲置,GPU 反复读写模型权重、KV 缓存,显存带宽拖慢速度。 优化方向:量化压缩、GQA/MQA 缩减 KV 缓存、分页注意力 PagedAttention、连续批处理。

算术强度判断瓶颈

算术强度 = 每字节显存传输对应的浮点运算次数:

  • 算术强度高 → 计算瓶颈;

  • 算术强度低 → 显存带宽瓶颈。

单用户对话 Decode 算术强度极低,绝大多数线上推理瓶颈是显存带宽,优先做显存压缩优化。


八、全文核心总结

  1. 训练 vs 推理:训练计算密集、推理显存密集,优化逻辑完全不同;

  2. 推理双阶段:Prefill 并行处理输入(计算瓶颈),Decode 逐字生成(显存带宽瓶颈);

  3. KV Cache 核心价值:缓存历史 K、V 向量,消除平方级重复计算,复杂度 O (N²)→O (N),推理必备;

  4. KV Cache 痛点:长序列、多并发大量占用显存,70B 模型 4k 序列缓存可达 12.8GB;

  5. 缓存优化方案:GQA 是当前最优折中,通过分组共享 KV 头大幅降低显存开销;

  6. 性能三大指标:TTFT 首字延迟、TPS 生成速度、端到端总延迟;

  7. 优化核心思路:Decode 阶段优先降低显存读写量(量化、GQA、分页 KV 缓存)。

本周学习路线

Day1 ✅ 推理基础 & KV Cache(本文) Day2 ⏭ 量化技术:INT8/INT4 模型压缩原理(GPTQ、AWQ、GGUF) Day3 ⏭ 主流推理框架实战:vLLM、TGI Day4 ⏭ 知识蒸馏:大模型轻量化方案 Day5 ⏭ 高并发模型服务化部署 Day6 ⏭ 进阶加速:推测解码 Speculative Decoding Day7 ⏭ 综合推理优化项目实战


下篇预告

KV Cache 解决重复计算,但显存资源依旧紧张。下一篇讲解量化技术,把 FP16 模型压缩为 INT8/INT4,显存占用直接缩减一半至 3/4,同步提速,详解主流量化算法落地细节。

点赞收藏不迷路,有推理部署相关疑问欢迎评论交流!

CSDN 配套标签

大模型推理KV CacheGQAMQA推理优化LLM部署AI大模型Transformer深度学习LLaMAvLLM显存优化

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐