热门搜索:和平精英 原神 街篮2 

您的位置:首页 > > 教程攻略 > ai资讯 >前女友面试官:大模型内存占用机制是怎样的?

前女友面试官:大模型内存占用机制是怎样的?

来源:互联网 更新时间:2026-08-25 14:33

大模型时代,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数据:

随机生成的 BF16 精度数据

计算公式为:

步骤拆解:

  1. 符号位 Sign = 1,表示负数。
  2. 指数位 Exponent = 17,计算:
  3. 小数位 Mantissa = 3,计算:

最终结果:

将三部分相乘,得到:-8.004646331359449e-34。

注意事项:

指数位全0和全1是特殊情况,不能套用上述公式。

02 全参训练和推理的显存分析

搞清楚了数据精度,相当于知道了不同零件的大小。但要估算整个生产线的资源需求,还得了解整个流程。接下来以最常见的混合精度训练为例,看看显存都去哪了。

混合精度训练

原理介绍

混合精度训练,就是把不同精度的数据类型混在一起训练。《MIXED PRECISION TRAINING》这篇论文采用了FP16和FP32混合,优化器使用Adam,流程如下:

MIXED PRECISION TRAINING 论文里的训练流程图

按训练逻辑梳理:

  • Step1:

    优化器先备份一份FP32精度的模型权重,并初始化FP32精度的一阶和二阶动量。
  • Step2:

    开辟新空间,将FP32的模型权重转换为FP16精度。
  • Step3:

    运行前向和反向传播,产生的梯度和激活值都用FP16精度存储。
  • Step4:

    优化器利用FP16的梯度和FP32的动,去更新备份的FP32模型权重。
  • Step5:

    重复Step2到Step4,直到模型收敛。

训练过程中,显存主要消耗在四个部分:

  • 模型权重本身(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

语言:简体中文

平台:互联网

游戏下载

热门手游

手机号码测吉凶
本站所有软件,都由网友上传,如有侵犯你的版权,请发邮件haolingcc@hotmail.com 联系删除。 版权所有 Copyright@2012-2013 haoling.cc