← 返回题目列表

如何估算大模型训练的算力和显存?6ND 法则是什么?

高频 困难 第 11 / 25 题 更新于 2026/07/28
训练算力FLOPs6ND显存估算显存构成

简化版

训练算力有个好用的经验公式:总计算量 C ≈ 6 × N × D(N 是参数量,D 是训练 token 数,单位 FLOPs)。由来是每个 token 每个参数,前向约 2 次浮点运算、反向约 4 次,合计约 6。显存则主要由四块构成:参数 + 梯度 + 优化器状态 + 激活。用混合精度 + Adam 时,「模型状态」约 16 字节/参数(FP16 参数 2 + FP16 梯度 2 + FP32 副本 4 + Adam 的 m/v 各 4),再加上随 batch 和序列长度增长的激活。掌握这两个估算,就能快速判断「训这个模型要多少算力、多少卡」。

详细版

算力估算:C ≈ 6ND

C(训练总 FLOPs) ≈ 6 × N(参数量) × D(训练 token 数)
  前向 ≈ 2ND
  反向 ≈ 4ND(反向约为前向的 2 倍)
  合计 ≈ 6ND
  • 例:训练 Llama-7B(N=7e9)用 2e12 token → C ≈ 6 × 7e9 × 2e12 ≈ 8.4e22 FLOPs。
  • 有了 C,除以「GPU 有效算力 × 利用率(MFU,常 40%~50%)× 卡数」就能估训练时间。

显存构成:四块

组成大小(混合精度 + Adam)说明
参数(FP16)2N 字节计算用的低精度权重
梯度(FP16)2N 字节与参数等大
优化器状态12N 字节FP32 参数副本 4 + m 4 + v 4
模型状态小计≈ 16N 字节上面三项之和
激活∝ batch × 序列长 × 隐藏维 × 层数随输入规模变化,长序列时很大
  • 模型状态 ≈ 16 字节/参数:7B 模型光这部分就约 112GB,远超单卡,必须靠 ZeRO/并行分摊。
  • 激活是另一大头,靠梯度检查点压缩。

MFU(模型算力利用率)

  • 实际训练达不到 GPU 理论峰值算力,MFU = 实际有效算力 / 理论峰值,通常 40%~55%。估算时间时要乘上它。

完整版教学

一、为什么要会”估算”

面试和实际工程都常问:「训一个 X B 的模型、Y T token,要多少卡、多久、多少显存?」能快速心算,说明你理解训练成本的构成。这背后就两把尺子:算力尺(6ND)显存尺(16N + 激活)。掌握它们,就能把「训练可行性」从拍脑袋变成算账。

二、6ND 从哪来:每 token 每参数约 6 FLOPs

核心事实:Transformer 里主要计算是矩阵乘,一次「乘加」算 2 个浮点运算。对每个 token、每个参数:

  • 前向:每个参数参与一次乘加,约 2 FLOPs/参数/token,全程 ≈ 2ND。
  • 反向:要算「对输入的梯度」和「对权重的梯度」两部分,计算量约是前向的 2 倍,≈ 4ND
  • 合计 ≈ 2ND + 4ND = 6ND

所以训练总算力 C ≈ 6ND。这个公式忽略了注意力里与序列长度平方相关的部分(在参数量主导时它占比不大)、以及一些逐元素操作,但作为数量级估算非常准,是业界通用的心算工具。它也正是 Scaling Law 里 C ≈ 6ND 的来源。

记忆钩子:前向 2、反向 4、合计 6——每 token 每参数约 6 FLOPs,训练算力 C≈6ND。

三、用 6ND 估训练时间

有了总算力 C,估时间只需再引入两个量:

训练时间 ≈ C / (卡数 × 单卡有效算力)
单卡有效算力 = 单卡理论峰值 FLOPS × MFU

MFU(Model FLOPs Utilization,模型算力利用率) 是关键的现实折扣:由于访存、通信、流水线空泡等开销,实际达到的算力远低于 GPU 标称峰值,通常只有 40%~55%。所以估算务必乘上 MFU,否则会严重低估时间。

举例:C ≈ 8.4e22 FLOPs,用 1024 张 H100(单卡 BF16 峰值约 1e15 FLOPS),MFU 取 40%:

有效总算力 ≈ 1024 × 1e15 × 0.4 ≈ 4.1e17 FLOPS
时间 ≈ 8.4e22 / 4.1e17 ≈ 2e5 秒 ≈ 2.4 天

数量级立现。这就是「多少卡训多久」的快速估法。

四、显存为什么是 16 字节/参数

训练显存的「模型状态」部分(不含激活)在混合精度 + Adam 下,逐参数账本是:

FP16 参数        2 字节   (前向反向计算用)
FP16 梯度        2 字节   (与参数等大)
FP32 参数副本    4 字节   (master weights,保证更新精度)
Adam 一阶动量 m  4 字节   (FP32)
Adam 二阶动量 v  4 字节   (FP32)
------------------------------------
合计            16 字节 / 参数

所以一个 7B 模型,光模型状态就 7e9 × 16 ≈ 112GB,一张 80GB 的卡都放不下。这直接解释了两件事:

  • 为什么必须 ZeRO/FSDP:把这 16N 分片到多卡,每卡只存 1/N。
  • 为什么优化器状态是分片首选:它占了 12/16(m + v + FP32 副本),是最大冗余,ZeRO-1 先分它。

五、别忘了激活:另一大头

模型状态(16N)是「随参数量固定」的部分,而激活显存随输入规模变化激活 ∝ batch × 序列长度 × 隐藏维 × 层数。在大 batch、长上下文时,激活能轻松超过模型状态,成为瓶颈。

应对手段(对应前面几道题):

  • 梯度检查点:少存激活、反向重算,把激活显存降到约 O(√层数)。
  • 序列/上下文并行:把长序列切到多卡。
  • 减小 micro-batch:直接降激活量。

完整的显存账 = 模型状态(16N,靠 ZeRO 分片)+ 激活(靠梯度检查点压缩)+ 一些临时缓冲。会拆这三块,就能定位「显存爆在哪、该开哪个优化」。

六、面试拆解算例

预训练题最好落到预算账和稳定性账。假设训练一个 7B 模型,目标 token 数是 1T,按常见粗估 6 × 参数量 × token 数,训练计算量约为 6 × 7e9 × 1e12 = 4.2e22 FLOPs。如果有效集群算力是 1e18 FLOPs/s,理想情况下也要约 42,000 秒;现实还要扣通信、数据加载、checkpoint 和故障恢复的损耗。

training_flops ≈ 6 * N * D
N = 7e9 parameters
D = 1e12 tokens
training_flops ≈ 4.2e22
账本关键变量常见瓶颈排查信号
数据账token 数、重复率、质量分脏数据和污染eval 异常偏高
算力账GPU 数、利用率、通信MFU 低step time 抖动
显存账batch、序列长、优化器状态OOM激活占用过高
稳定性账学习率、精度、梯度loss spikeoverflow/NaN
语料 -> 清洗去重 -> tokenization -> 分布式训练 -> checkpoint -> 评测
  |        |             |                |             |
 质量     覆盖率         吞吐              可恢复         能力验证

所以回答「如何估算大模型训练的算力和显存?6ND 法则是什么?」时,不能只说某个技巧“省显存”或“加速”。要说明它省的是哪一笔账、牺牲了什么、线上训练日志里应该观察哪个信号。

七、常见误区与追问

  • 误区:训练算力就是前向的量。 反向约是前向 2 倍,合计 6ND,不是 2ND。
  • 误区:显存主要是参数占的。 参数只占 2N,优化器状态 12N 才是大头,加上激活;模型状态共约 16N。
  • 追问:6ND 里 6 怎么来的? 前向 2 + 反向 4(反向算输入梯度和权重梯度,约 2 倍前向)。
  • 追问:为什么估时间要乘 MFU? 实际算力因访存/通信/空泡只有峰值的 40%~55%,不乘会严重低估。
  • 追问:16 字节/参数怎么拆? FP16 参数 2 + FP16 梯度 2 + FP32 副本 4 + Adam m 4 + v 4。
  • 追问:纯 BF16、不存 FP32 副本能省吗? 可以更省,但通常保留 FP32 主权重保证收敛;也有更激进的低比特优化器状态方案。
  • 追问:激活为什么随序列长度增长? 每层每个 token 位置都要缓存中间结果,长序列 + 大 batch 时激活线性放大。

八、加强记忆

两把尺子钉死:算力 C ≈ 6ND(前向 2 + 反向 4,每 token 每参数约 6 FLOPs),估时间再除以「卡数 × 单卡峰值 × MFU(40%~55%)」显存 = 模型状态 ≈ 16N(FP16 参数 2 + FP16 梯度 2 + FP32 副本 4 + Adam m 4 + v 4,优化器占大头 12N)+ 激活(∝ batch × 序列长 × 隐藏维 × 层数)。由此秒懂两件事:16N 撑爆单卡 → 必须 ZeRO 分片(先分优化器状态);激活撑爆长序列 → 开梯度检查点。记住「6ND」和「16 字节/参数」,训练成本就能随手估出数量级。