前女友面试官:大模型内存占用机制是怎样的?
大模型时代,GPU显存是实实在在的硬通货。能用好它,训练和推理的效率往往能差出一大截。
这篇文章会围绕单卡场景,把大模型的内存占用机制彻底讲清楚。理解了这个,后续不管是做训练还是推理,心里都会更有底。
主要会回答三个核心问题:
- 给你一个模型参数量,怎么估算训练和推理时的显存占用?
- Lora相比全参训练,到底省的是哪部分显存?Qlora相比Lora,又省了哪部分?
- 混合精度训练的具体流程是怎样的?
这些内容也是面试中的高频考点。梳理这些知识,既是巩固,也希望能帮到正在准备秋招或相关工作的小伙伴。
这篇文章会聚焦于单卡训练或推理时的显存占用,来做一次系统性的分析。部分知识点点到为止(毕竟有些细节我也还没完全吃透),但尽力保证整篇文章逻辑流畅、通俗易懂。
01 数据精度
计算显存,最底层要搞清楚数据精度,它直接决定了一个数据占多大空间。
基础换算关系:
1 byte = 8 bits
1 KB = 1,024 bytes
1 MB = 1,024 KB
1 GB = 1,024 MB
举个例子,一个包含10亿参数的模型,如果每个参数用32bit(4byte)存储,直接加载就需要占用4GB的显存。
常见精度类型
掌握下面这几种常见的精度就足够了,其他的可以触类旁通。图片源自英伟达安培架构白皮书:
各种精度的数据结构
从图里可以看到,浮点数由三部分组成:符号位、指数位和小数位。符号位固定为1位(0正1负),指数位决定了浮点数的表示范围,小数位决定了精度。
注意,TF32虽然有“32”这个名字,但实际只有19bit。BF16(Brain Float 16)由Google Brain团队提出。
具体计算例子
抽象的概念讲再多,不如一个具体的例子来得直接。下面用BF16来演示,如何通过符号位、指数位和小数位计算出最终数值。
下面是随机生成的一个BF16数据:
计算公式为:

步骤拆解:
- 符号位 Sign = 1,表示负数。
- 指数位 Exponent = 17,计算:
。 - 小数位 Mantissa = 3,计算:
。
最终结果:
注意事项:
02 全参训练和推理的显存分析
搞清楚了数据精度,相当于知道了不同零件的大小。但要估算整个生产线的资源需求,还得了解整个流程。接下来以最常见的混合精度训练为例,看看显存都去哪了。
混合精度训练
原理介绍
混合精度训练,就是把不同精度的数据类型混在一起训练。《MIXED PRECISION TRAINING》这篇论文采用了FP16和FP32混合,优化器使用Adam,流程如下:

MIXED PRECISION TRAINING 论文里的训练流程图
按训练逻辑梳理:
- 优化器先备份一份FP32精度的模型权重,并初始化FP32精度的一阶和二阶动量。
Step1:
- 开辟新空间,将FP32的模型权重转换为FP16精度。
Step2:
- 运行前向和反向传播,产生的梯度和激活值都用FP16精度存储。
Step3:
- 优化器利用FP16的梯度和FP32的动,去更新备份的FP32模型权重。
Step4:
- 重复Step2到Step4,直到模型收敛。
Step5:
训练过程中,显存主要消耗在四个部分:
- 模型权重本身(FP32+FP16)
- 梯度(FP16)
- 优化器(FP32)
- 激活值(FP16)
三个小问题
第一个问题:为什么不全部用FP16?那样计算更快,显存占用更少。
答案在于FP16的精度范围远窄于FP32,这会引发数据溢出和舍入误差,导致梯度消失,训练无法进行。所以必须依赖FP32来保证精度。不过,现在很多训练改用BF16,它范围更宽,至少不会出现数据溢出,业界实践也证明,大模型对数值范围的需求优先级高于精度。
第二个问题:为什么只对激活值和梯度做半精度优化,却新增了一个FP32的模型副本?这样显存不会更大吗?
答案是不会。激活值的占用与batch_size和序列长度强相关,在实际训练中,激活值往往是显存消耗的大头。对激活值进行正向优化带来的节省,远大于备份模型参数的额外开销,最终显存是减少的。
第三个问题:显存和内存一样,有静态和动态之分。上面提到的哪些是静态,哪些是动态?
通常的划分:
- 优化器状态、模型参数
静态:
- 激活值、梯度值
动态:
因此,很难精确计算实际运行时的显存峰值。面试时,可以忽略激活值的计算,并把梯度当作静态考虑。

动态监控显存图
来个小测试
现在理论说得差不多了,来实操一下。对于llama3.1 8B模型,用FP32和BF16混合精度训练,采用AdamW优化器,模型训练时占用显存大概是多少?
解:
- BF16 (16G) + FP32 (32G) = 48G
模型参数:
- BF16 (16G) = 16G
梯度参数:
- FP32 (32G) + FP32 (32G) = 64G
优化器参数:
- 48G + 16G + 64G = 128G
不考虑激活值,总显存:
推理与KV Cache
原理理解
推理阶段,显存主要花在模型参数本身,以及现在广泛使用的KV Cache上。
KV Cache不是为了省显存,而是为了降低延迟,用显存换速度。
具体来说,推理本质上是不断重复“生成下一个token”的任务。生成当前token,只依赖当前的QKV和之前所有KV。因此,可以维护并不断更新这个KV,避免重复计算。
KV Cache 动态实现
一个常见疑问:为什么没有Q Cache?因为生成当前token只依赖当前的Q,这是由Self-Attention的公式决定的。
公式中,在序列的第t行,只与前面的K和V有关,这意味着不需要保存每一步的Q。更本质地说,矩阵乘法的数学特性决定了这一点。
计算KV Cache显存
KV Cache显存的计算公式如下:
公式中的4个参数相乘,代表KV在模型每一层所有隐藏向量的总和。第一个2指K和V两部分,第二个2对应半精度的字节数。
以llama7B为例(hiddensize=4096, seqlength=2048, batchsize=64, layers=32),计算结果是68G。
可以看到,在大批量、长句子的场景下,KV Cache的显存占用相当可观。但如果是单batch,KV Cache大约只占1G,约为模型参数显存的一半。
MQA和GQA
如果觉得KV Cache占用的显存还是太多,MQA和GQA就是用来进一步压缩的方法。目前主流大模型基本都采用了这些技术。
三种 KV 处理方式
方法不难理解,核心在于共享多头的KV,这是一个很朴素的剪枝思路。最左侧是基础的MHA(多头自注意力),中间是GQA(分组查询注意力),保留了几组KV头;右侧是MQA(多查询注意力),只保留1组KV头。目前GQA用得更多,在降低显存、提升速度的同时,性能损失更小。
MHA的KV Cache计算公式为:
有两个额外注意点:一是MQA和GQA模型可以从头开始训练,也可以像相关论文那样,基于开源模型修改结构后继续预训练。目前大多从头训练,以保证训练和推理的模型结构一致。
03 Lora和Qlora显存分析
前面详细分析了全参微调训练和推理的显存。一个很现实的问题是:现在主流都是PEFT(高效参数微调),全参训练的资源要求太高;推理阶段也需要量化。那这些场景下的显存如何分析?
理解了前两章,再来看这些,会轻松很多。显存分析的核心,就是理清流程和数据精度,分析方法是一样的。接下来详细拆解Lora和Qlora的显存占用。
Lora
Lora的原理不算复杂:在原始权重矩阵旁路新建一对低秩的可训练权重。训练时只更新旁路,极大减少了训练参数量(从d*d降为2*d*r)。
Lora 原理图
借用一下前面全参训练的分析思路,设定为BF16模型、AdamW优化器、Lora参数也是BF16,设定1字节模型参数对应的显存为φ。
首先是模型权重本身。需要加载原始模型和Lora旁路模型,Lora部分占比不到2个数量级,可以忽略。因此显存占用约为2φ。
然后是优化器部分。优化器只针对需要更新的参数,即Lora模型权重。同样,因为数量级太小,可以忽略,占用显存约为0φ。
最让人困惑的是梯度部分。有观点说原始模型也要参与反向传播,所以需要梯度;也有观点说原始模型不更新,所以只需要Lora部分的梯度。正确答案是:不需要计算原始模型部分的梯度,基本不占用显存。因此梯度部分显存也近似为0φ。
综上,不考虑激活值,Lora微调训练的显存占用约为2φ。一个7B模型用Lora训练,大概需要14G显存。
可以验证一下。LlamaFactory给出的训练任务显存预估表格,7B模型Lora训练的显存消耗与我们估算的接近,同时也符合之前对全参、混合精度训练的显存分析。
Llama Factory 的表格
QLora
QLora,全称量化Lora,是Lora之后又一个广泛用于大模型PEFT的方法。核心思路是进一步压缩模型精度,然后用Lora训练。理解起来不难,但细节不少。
QLora的整体思路
QLora出自论文《QLORA: Efficient Finetuning of Quantized LLMs》。论文的核心是一种新的量化方法,重点在量化,而非Lora。
有些人不了解,以为量化Lora是对Lora部分参数进行量化,因为只有Lora参数参与训练。但理解上面Lora的朋友就能明白,实际上原始模型虽然不更新参数,但仍需参与前向和反向传播。QLora优化的是Lora里显存占大头的模型参数本身。
那么,QLora是把原始模型参数从16bit压缩到4bit,然后更新这个4bit参数吗?注意,这里要区分“计算参数”和“存储参数”。计算参数是在前后向传播中参与实际计算的参数;存储参数是不参与计算、一开始加载的原始参数。
QLora的做法是:先将16bit的原始模型参数加载并量化为4bit,作为存储参数。在需要计算时,再将这4bit参数反量化为16bit,作为计算参数,用完即释放。也就是说,QLora训练时所有数据的精度都与Lora一样,只是加载的模型是4bit,计算时会反量化到16bit。
而Lora部分的参数全程是16bit,不需要量化。
这比Lora多了一步量化和反量化的操作,训练时间自然会变长。一般来说,QLora训练比Lora要多用约30%的时间。
QLora的技术细节
QLora主要有三个创新点:
- 传统量化假设参数均匀分布,而NF4基于参数正态分布的假设,大幅提升了量化精度。
NF4量化:
- 对第一次量化后用于反量化的锚点参数,再进行一次量化,进一步降低显存。
双重量化:
- 为防止OOM,在GPU显存紧张时,可将参数临时转移到CPU内存。
优化器分页:
显存分析
理解了QLora的运行思路,显存占用部分就很清晰了。QLora的主要显存消耗在于4Bit量化后的模型本身,即0.5φ。这里同样没有考虑Lora部分的参数和量化计算中可能的额外显存。
回顾之前的表格,这个估算也基本符合预期。
最后,用一张表格总结前面所有的显存分析:
来源:https://zhuanlan.zhihu.com/p/713256008
-
- 关于宇宙的好的网名有哪些
- 角色扮演 | 1
- 网名