ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

大模型训练显存估算与混合精度实战指南

大模型训练显存估算与混合精度实战指南 跑大模型训练的人谁没被 “CUDA out of memory” 那行红字折磨过大模型训练做到中后期显存就不只是“够不够用”的问题而是必须带着计算器去预判的问题。这篇是这个系列的第三篇我把大模型训练的显存估计和混合精度训练放在一起聊因为它们在实战中几乎总是绑在一起出现显存不够用的时候第一反应不是换卡而是先看看有没有把精度浪费在不该浪费的地方再考虑怎么用混合精度把空间省出来。这篇文章先给出一套能直接套用的显存估算方法把 7B、13B 这种常见规模算给你看再讲清 FP32、FP16、BF16 在训练中的真实差异以及大模型训练混合精度为什么既能省显存又能提速最后是一份带代码的实操清单和踩坑记录。不管你正在准备训自己的模型还是单纯想搞清楚为什么隔壁团队能在一张 A100 上跑起来而你不能这篇都值得存一下。1. 训练显存不是只有模型权重1.1 静态开销权重、梯度、优化器状态三件套绝大多数人对显存的直观理解就是“模型多大显存就多大”。这个直觉在推理场景下基本成立但在训练时误差很大。训练时显存实际上由四个部分构成模型权重、梯度、优化器状态、激活值。前三个属于静态开销基本不随训练轮次变化最后一个是动态大头会在前向过程中持续波动。先说静态三件套。模型权重必须常驻显存这个不用解释。梯度呢反向传播逐层算出梯度后并不是马上丢掉的优化器要做整体更新分布式训练里要跨卡聚合梯度在权重更新之前必须完整保存在显存里。梯度的数据量和模型权重完全一样相当于把模型本体再存了一份。优化器状态则是很多人容易漏掉的一项以最常用的 AdamW 为例它会为每个参数维护两个状态一阶动量momentum和二阶动量variance而且这两个状态通常都保留为 FP32 精度。也就是说一个 7B 模型用 AdamW 训练光是优化器状态就相当于把 7B 个 FP32 数存了两次算下来比模型本身还大。给你一个厨房备菜的类比权重是你灶台上正在用的食材梯度是切好还没下锅的菜码优化器状态是贴在墙上的配方笔记三者缺一不可但只有真正做完菜你才会发现最占台面的其实是临时堆在那里等着下锅的备菜——那就是激活值。1.2 激活值被忽略的显存大头激活值activation指的是前向传播过程中每一层算出来的中间张量。Transformer 的每一层要做自注意力、LayerNorm、MLP 等一连串计算反向传播时需要这些中间结果才能回传梯度所以不能算完就扔。层数越深、序列越长、batch 越大激活值的体积就滚得越大。很多人的 OOM 不是发生在加载模型的时候而是训练跑到第几百步才突然炸掉原因就在这里加载阶段只有静态三件套训练开始后激活值叠加上来峰值一冒显存就崩了。明白这个结构之后你拿到一个新模型就不该再笼统地问“这个模型要多少显存”而应该拆成两个问题静态三件套占多少激活值峰值占多少。这两个问题的估算方法就是下一章要说的重点。2. 显存估算先算静态再算动态2.1 一张表看懂精度与字节数在估算之前先记住不同数据类型对应的字节数。大模型训练里最常见的精度就是 FP32、FP16、BF16偶尔还会见到 INT8 精度的优化器。数据类型字节数指数位尾数位训练中的典型用途FP324823优化器状态、主权重FP162510旧卡上的混合精度前向计算BF16287新卡上的混合精度前向计算INT81无无量化优化器状态一个参数占用几个字节拿参数量 P 一乘就出来了。假设参数量是 700 亿也就是 P70×10^9那么 FP32 下光权重就是 70×4280GB看到这个数字你马上就能理解为什么 70B 模型训练几乎不可能用纯 FP32 单卡完成。2.2 常见训练组合的系数表大模型训练里最常见的优化器组合就那么几种我直接把“每参数字节数”的系数列出来。这里的 P 表示参数量单位是字节。训练配置权重梯度优化器状态每参数总字节7B 模型静态显存FP32 SGD4P4P08P约 56GBFP32 AdamW4P4P8P16P约 112GB混合精度 AdamW2P2P8P12P约 84GB有点反直觉的是混合精度训练虽然把权重和梯度降到了 FP16/BF16但优化器状态依然是 FP32 的两份所以总字节系数是 22812P而不是单纯的一半。这也是为什么大家普普通通说“7B 混合精度训练需要 80 多 GB 显存”的来源。顺便说一句纯 FP32 的 AdamW 需要 16P一个 7B 模型就吃掉 112GB单卡基本没戏而混合精度把权重和梯度各减半降到 84GB配合激活值控制和梯度检查点就有机会塞进单张 80GB 的卡里。省下来的这 28GB就是混合精度的最大意义之一。2.3 手把手算一次 7B、13B、70B套用 12P 这个系数你拿计算器直接乘就行7B 模型12 × 7 84GB不含激活值13B 模型12 × 13 156GB不含激活值70B 模型12 × 70 840GB不含激活值。看到 840GB 不要慌这是单卡全量放置的数字实际训练 70B 都会走多卡并行加 ZeRO 分片把权重、梯度、优化器状态平均拆到每张卡上。后面第 5 章会讲怎么拆。还要特别注意这里的数字只是静态开销。一个 7B 模型的 84GB 算出来如果手里只有一张 80GB 的 A100听起来勉强够但一旦前向传播的激活值冲上来100GB 照样瞬间爆掉。所以估算永远是“先算静态再给动态留余量”我自己的习惯是按照公式结果的 1.2 到 1.3 倍来选卡宁可富余别赌运气。2.4 激活值怎么估经验公式加实测激活值没有静态三件套那么规整但可以按 Transformer 的常见实现估一个量级。经验上每个 Transformer 层要保存的中间激活元素数大致是$$batch \times seq_len \times hidden_size \times k$$其中 k 是一个和实现相关的经验系数常见实现不开梯度检查点时大约在 30 到 40 之间按 2 字节存储。以 Llama-7B 级别的配置为例hidden_size 取 4096层数取 32序列长度 2048batch_size 为 1带入后$$1 \times 2048 \times 4096 \times 32 \times 34 \times 2 \approx 18GB$$也就是说单序列长度 2048 的 7B 模型光激活值峰值就逼近 18GB。一旦 batch 开到 4这一项就变成 72GB整张 80GB 的卡直接就没有余量了。这还只是保守估计实际因为 dropout mask、临时变量、注意力矩阵的实现差异峰值可能更高。这个数字也解释了一个现象为什么大模型训练里大家那么怕“长序列”sequence length 一翻倍激活值和序列长度的平方项一起涨对显存的杀伤力远比增加 hidden size 要大。所以做显存预算的时候激活值必须当成头号变量来对待。最可靠的方法是本地跑一个小配置用 PyTorch 的 torch.cuda.max_memory_allocated() 直接量出真实峰值经验公式只能用来做出发前的预判。3. 混合精度训练的原理与选型3.1 FP16 和 BF16 到底差在哪混合精度训练的核心思路是把前向计算和梯度计算中那些对精度不太敏感的算子从 FP32 降到半精度从而省显存、降带宽压力、利用 Tensor Core 加速。但半精度有两个兄弟FP16 和 BF16长相相近脾气完全不同。FP16 是 1 位符号、5 位指数、10 位尾数最大能表示到 65504最小正规数大约是 6.1e-5。问题就出在这个范围上训练时反向传播的梯度经过层层乘法链式法则很多数值会掉到 1e-5 以下一旦小于最小正规数FP16 里存的就是 0。梯度变成 0意味着权重根本不会更新模型原地踏步。反过来碰到中间结果稍大一点超过 65504 就变成无穷大直接炸掉整个训练。上下两个方向都容易出问题这是 FP16 最麻烦的地方。BF16 是 1 位符号、8 位指数、7 位尾数。指数位和 FP32 一样多所以它的表示范围和 FP32 几乎一致极小值到极大值都能覆盖天然不会出现 FP16 那种动不动溢出、下溢的情况。代价是尾数只有 7 位十进制有效数字大约只有 2 到 3 位精度比 FP16 还低。这就带来一个很有意思的局面BF16 范
返回列表