怎么考「微调 7B 模型,全参和 LoRA 各要多少显存?怎么估算?

Q4-06训练资源估算常见

谁在问:算法工程向二面(经典压轴);基建/训练平台团队;面试官会临场换规模换精度逼你现推

开场怎么问

微调一个 7B 模型大概要多少显存?全参和 LoRA 分别说一下。

换个问法

  • 这个数你是怎么算出来的?如果换成 fp32 优化器呢?换成 70B 呢?
  • 你说 LoRA 只要 20 G,那 20 G 里面都是什么?

五层追问链

左边照着问,右边对着听。最后一层是压力面,不必每个候选人都问到。

12 和 16 到底差在哪?什么时候是 12,什么时候是 16?

期望
差的是 fp32 主权重副本那 4 字节。经典混合精度训练为了数值稳定会保一份 fp32 主权重,参数更新在 fp32 上做,再转成 BF16 用于前向——这是 16。现代 BF16 训练(BF16 动态范围与 fp32 相同,不像 fp16 那样需要 loss scaling)可以不保主权重,直接在 BF16 权重上更新——这是 12。具体取决于框架和优化器配置。
信号
能说出"差的是 fp32 主权重"——真读过或真配过。答"看资料写的不一样"——背的。能补一句"BF16 的动态范围和 fp32 相同,所以可以省掉主权重,fp16 不行"——非常强的信号。

LoRA 的可训参数量怎么算?给个 7B 上 r=16 的数。

期望
r×(d+k) 逐矩阵累加。r=16、all-linear、7B ≈ 40 M(是 r=64 那 160 M 的四分之一,线性关系)。三件套约 0.5 G。能当场推的人,说明自己配过而不是抄过。
信号
能给出"和 r 成线性"这个关系并现场折算——理解了公式。答"很少,可以忽略"——结论对但没算过。

激活显存怎么估?怎么压?压了有什么代价?

期望
量级上正比于 batch × seq_len × hidden × 层数(注意力部分对 seq_len 还有更强的依赖)。压的手段按收益排序:缩短序列长度(收益最大)> 梯度检查点(慢 20–30%)> 减小 batch 用梯度累积补 > FlashAttention 这类算子级优化。代价:前三个都是拿时间换显存。
信号
能给出"缩短序列长度收益最大"这个优先级——真优化过。只会答"开梯度检查点"——知道招式但没有排序。

换成 70B 呢?单卡放不下的时候怎么切?

期望
70B 全参 840 G 起,必须多卡:ZeRO-1 切优化器状态 / ZeRO-2 再切梯度 / ZeRO-3 连参数一起切(通信量依次上升),必要时加 CPU offload;或用 FSDP、张量并行。70B LoRA 约 140 G 基座 → 四卡;70B QLoRA 约 35 G → 单卡 80 G 可行
信号
知道 ZeRO 三个 stage 各切什么——有分布式训练概念。能说出"stage 越高省得越多但通信越重"——真跑过多卡。

你按这套账估了 20 G,申请了一张 24 G 的卡,实际一跑就 OOM。给你半天时间,你怎么排查、怎么改,什么时候你会去要更大的卡?

期望
  • 先定位是哪一块超了,别乱调:把 batch 降到 1、序列长度降到 512 再跑一次。如果这样还 OOM,说明是基座/参数那块算错了(比如 all-linear 下可训参数比预期多、或者框架偷偷保了 fp32 主权重);如果这样不 OOM,说明超的是激活,问题在 batch 和序列长度上。这一刀是本题的核心动作。
  • 常见的真实原因max_seq_len 用了默认 4096 而数据实际只有 500;没开梯度检查点;框架默认 fp32 优化器;显存碎片(长序列变长时反复分配);评估阶段的 batch 没单独设置,训练能跑评估就炸。
  • 修复顺序:按真实数据的长度分位数(比如 P95)设 max_seq_len → 开梯度检查点 → batch=1 + 梯度累积 → 确认优化器精度配置 → 最后才考虑 QLoRA。
  • 什么时候去要卡:把上面全做完仍然差得远(比如还差一倍以上),说明是方案层面的问题——这时候要么换更小的基座,要么申请资源,但申请时要带着这份三块账去谈,而不是"感觉不够"。
  • 加分:主动提"先看 nvidia-smi 的峰值出现在训练的哪个阶段"——峰值出现在反向传播说明是激活,出现在优化器 step 说明是优化器状态。
信号
  • 直接开始降 batch 试——能解决问题但没有定位过程,说明靠试不靠算。
  • 能提出"batch=1 + 短序列跑一次来切开参数侧和激活侧"——这是本题最强信号,它复用了本项目一贯的排障骨架:先用一个极端配置把两侧切开,再逐步加回来。
  • 能说出"评估阶段 batch 没单独设"这类具体坑——真被坑过。
  • 一上来就说"要更大的卡"——没有排查能力。

危险信号

听到这些话,基本可以判定是背题而不是做过。

"7B 全参大概 80 G"只有结论没有账本。换成 13B 或换优化器就答不上来。
全程不提激活本题最主要的判负点。 说明没跑过训练,因为真跑过的人都被激活坑过。
"LoRA 显存很少,几个 G 就够"忘了基座那 14 G 还在。LoRA 省的是三件套不是基座。
"QLoRA 省的是优化器状态"位置错误。省的是基座权重。
"梯度和权重都是 fp32"混合精度训练的基本概念没有。
"开梯度检查点就不占显存了"只是大幅减少激活,且有 20–30% 的时间代价。
报一个精确到小数点的数字但说不出构成含糊回答的典型形态:用精度冒充理解。 追问"这 19.4 G 里面都是什么"即崩。
OOM 了只会降 batch靠试不靠算,没有定位方法。

评分卡

  1. 能把显存拆成基座 / 可训参数三件套 / 激活三块
  2. 知道每参数 12~16 字节,并能说出差异来自 fp32 主权重
  3. 知道 LoRA 冻结基座后,三件套只按可训参数算
  4. 能现场估算 LoRA 可训参数量(r×(d+k) 逐层累加)
  5. 主动提到激活显存,并知道它随 batch 和序列长度变化
  6. 知道梯度检查点是拿时间换显存,代价约 20–30%
  7. (加分)能换算到 70B 并说出多卡切分方案
  8. (加分)知道缩短序列长度比降 r 更有效
  9. (加分)OOM 排查时能用极端配置切开参数侧与激活侧
参考答案与考察意图面试中途别看这一段

考察意图

这道题背不下来——面试官会临时换模型规模、换优化器、换精度、换 batch,逼你现场重推。它筛的是:

  1. 有没有一本账。能不能把显存拆成"基座 / 可训参数三件套 / 激活"三块分别算。
  2. 知不知道每参数多少字节,以及为什么。12 还是 16,差在哪,这是最能一句话拉开差距的地方。
  3. 记不记得激活。这是实际 OOM 的头号原因,也是最常被漏掉的一项。只报参数相关的显存、不提激活的候选人,基本可以判定没真跑过训练。

参考答案

图 2 · 60 分与 90 分差在哪:代价、演进、怎么验证

60

60 分答案(及格线)

全参微调一个 7B 模型,用 Adam 优化器、BF16 混合精度:

  • 权重 2 字节/参数 → 14 G
  • 梯度 2 字节/参数 → 14 G
  • 优化器状态(Adam 的两个 FP32 动量)8 字节/参数 → 56 G

加起来 84 G,再加上激活,大概 88 G 左右,单张 80 G 的卡放不下。

LoRA 的话基座冻结,只占权重那 14 G,可训参数很少所以三件套可以忽略,加激活大概 20 G 左右。QLoRA 把基座量化到 4 bit,14 G 变 3.5 G,总共不到 10 G。

账目基本对,及格。但没解释 12 字节这个数的来源、没说激活怎么估、没法应对"换成 fp32 主权重呢"这类追问。

90

90 分答案(有生产经验的回答)

我分三块算:基座权重、可训参数的三件套、激活。三种方案的差别几乎全在"梯度和优化器状态跟谁走"。

第一块和第二块——每参数多少字节:

经典混合精度口径(含 fp32 主权重副本,ZeRO 论文的算法):
  BF16 权重 2 + BF16 梯度 2 + FP32 主权重 4 + FP32 一阶动量 4 + FP32 二阶动量 4 = 16 bytes/param

现代纯 BF16 口径(不保 fp32 主权重):
  BF16 权重 2 + BF16 梯度 2 + FP32 双动量 8                              = 12 bytes/param

所以正确的答法是 "12 到 16 字节每参数,取决于优化器实现有没有留 fp32 主权重副本",7B 就是 84 到 112 G。报单一数字的人通常是背的。

LoRA 的关键在于:基座冻结 → 基座只占"权重"那一项(BF16 就是 2 字节/参数 = 14 G),没有梯度也没有优化器状态;那 12~16 字节的账只按可训参数算。

可训参数量是能算出来的,不用猜。LoRA 每个 d×k 矩阵新增 r×(d+k)。7B 典型结构(hidden 4096、intermediate 11008、32 层)、all-linear、r=64:

原文示意
每层:q,k,v,o  4 × 64×(4096+4096)  ≈ 2.10 M
      gate,up  2 × 64×(4096+11008) ≈ 1.93 M
      down     1 × 64×(11008+4096) ≈ 0.97 M   →  小计 ≈ 5.0 M
32 层 ≈ 160 M 可训参数(约占基座 2.3%),三件套 ≈ 160M × 12 ≈ 1.9 G

第三块——激活,也是最容易漏的一块: 它跟 batch × seq_len × hidden × 层数 走,是 2–20 G 的浮动项。"我算出来 20 G,结果 80 G 卡还是 OOM",九成是这里。 压它的标准手段是梯度检查点:前向只存少数中间结果、反向时重算,拿时间换显存,典型代价是训练慢 20–30%。

完整对照表(7B,r=64 all-linear):

基座 可训参数三件套 激活 合计
全参(12 B/param) 14 G 70 G 2–20 G 约 88 G
全参(16 B/param) 14 G 98 G 2–20 G 约 116 G
LoRA r=64 14 G 约 1.9 G 2–20 G 约 19 G
QLoRA r=64 3.5 G(NF4) 约 1.9 G 2–20 G 约 9 G

这套讲法的价值在于面试官换任何参数我都能重算。 比如换 70B:全参就是 70×12 = 840 G,八张 80 G 卡(640 G)都不够,必须 ZeRO-3 分片再加 offload;LoRA 是 140 G 基座,两张 80 G 勉强、实际要四张;QLoRA 是 35 G 基座,单张 80 G 卡就能跑——这也是 QLoRA 当年真正的意义所在。

攒够了去组卷页一键生成可打印的面试题单