训练一个 10B 参数的模型,你最心疼的是什么?

显卡。花十几万买来的 80GB 显存,一大半不是被模型本身吃掉的,而是被一堆看不见的东西吃掉的。
我最早也以为,训练时显存的大头是参数。直到我被 OOM 反复折磨,认真算了一笔账,才发现自己错得离谱:10B 参数,FP32 下模型只占 40GB。但 Adam 优化器会给每个参数挂两本“账本”——一阶动量 m 和二阶动量 v,各占 4 字节;梯度再占 4 字节。每个参数一共 16 字节,10B 参数就是 160GB。你手里的 80GB 显卡,连零头都装不下。
8-bit Adam 干的事,就是把那两本“账本”从 32-bit 压到 8-bit——账本从 8 字节缩到约 2 字节,整体显存砍掉大约四成。
听起来很简单?把 32-bit 数据压成 8-bit,谁不会。但这件事真正的难点是:你要是直接硬压,模型会训不出来,而且翻车的方式极具欺骗性——Loss 照常下降,指标却提不上去。这背后的坑,以及那个相当聪明的解法,才是这篇文章真正想讲的。
账本里记的到底是什么
Adam 每一步更新,要维护两个状态,它们是跟模型同形状的两个张量:
- m(一阶动量):梯度的指数滑动平均,相当于对梯度方向做了平滑,去掉噪声;
- v(二阶动量):梯度平方的指数滑动平均,代表每个参数最近“闹腾”得有多凶。
每轮更新长这样:
m = β1·m + (1-β1)·g # 梯度平滑
v = β2·v + (1-β2)·g² # 梯度强度
θ = θ - lr·m / (√v + ε) # 更新参数
v 是 Adam 自适应学习率的核心:某个参数梯度一直很大,v 就大,分母大,步长自动变小;反之步长变大。没有这两本账,Adam 就退化成带惯性的 SGD。
对 10B 模型来说,这两本账要占 80GB——比模型本身还大一倍。
直接量化?翻车的方式很隐蔽
在 8-bit Adam 之前,业内主流做法是:模型参数和梯度用 FP16,但两本账坚持用 FP32。理由是优化器状态对精度敏感,压不得。
Dettmers 团队上来先做了一件暴力的事:把 m 和 v 直接压成 8-bit,跑起来看。结果不出所料,性能掉了。但他们发现了一个更值得玩味的现象:掉性能不是因为 8-bit 精度不够,而是量化方式用错了。
问题出在 v 的分布上。v 里有一小撮极端离群值,比中位数大几百到几千倍。如果整个张量用一个全局 scale(按最大值)来量化,离群值会把区间撑爆——剩下 99.9% 的普通数值全被压缩到极少数档位里,几乎失去区分度。
打个比方:一群人里站着一个姚明,你按姚明的身高做一把尺去量所有人,结果全班同学的身高都约等于零。
解法:把“姚明”关进单独的房间
Dettmers 的思路直接得近乎朴素:既然离群值只有少数,就别让它们祸害全局。把整个状态张量切成若干块,每块 2048 个值,各自按块内最大值算自己的 scale。离群值再狂,也只能污染它所在的那一块,其它块完全不受影响。
We find that block-wise quantization, where each block of 2048 values is quantized independently, is crucial for the performance of 8-bit optimizers.
这个设计真正巧妙的地方,是把全局均匀精度换成了局部自适应精度。量化误差被精准地分配给了那些数值极大、对梯度方向影响有限的离群值;而大量中等数值所在的块,scale 小,量化反而更精细。误差被挤到了不重要的地方。
但这还不够。即使 block-wise,v 里偶尔还是会有极端值拖垮整个块。论文又加了一道保险——stable embedding:专门针对 v 的极端值,把超过阈值的少量离群值单独拎出来存成 16-bit,其余照常走 8-bit。这些特例通常不到全部值的 0.1%,代价可忽略,却换来训练稳定。
再加上动态量化——每一轮训练都根据当前数值范围重新计算 scale,而不是训前算一次就完事——三件套合体,就是 8-bit Adam 的全部秘密。
省了 75% 的账本,精度几乎没掉
论文在 GLUE 和 SuperGLUE 共 28 个任务、ImageNet、语音识别、机器翻译上做了对比,结论相当干净:
- 最终准确率与 32-bit Adam 基本持平,差距在 0.1% 以内;
- 优化器状态从每参数 8 字节降到约 2 字节,省了 75%;
- 整个训练流程的总显存平均减少 42.6%;
- 首次在单台机器上完成 175B 参数模型 的训练。
还有个反直觉的点:8-bit 训练有时比 32-bit 还快。状态变小了,内存带宽压力小了,cache 命中率反而更高。省显存的同时,顺手把速度也提了。
跟混合精度、ZeRO 到底是什么关系
这里有个我踩过的坑:我一度以为有了混合精度就不需要 8-bit Adam,后来发现完全不是一回事——它俩解决的是不同的问题:
| 方案 | 省的是什么 | 能跟 8-bit Adam 叠加吗 |
|---|---|---|
| FP16/BF16 混合精度 | 模型参数和梯度(4→2 字节) | 能,且常见 |
| 梯度检查点 | 前向传播的激活值 | 能 |
| ZeRO / FSDP | 把状态分片到多卡,总量没减 | 能,是黄金搭档 |
| 8-bit Adam | 优化器状态(8→约 2 字节) | — |
混合精度把参数和梯度减半,但账本还是 FP32;8-bit Adam 专砍账本;ZeRO 是把账本分片到多张卡上,每张卡的负担变小,但总量不变。三者叠加,才是今天训练超大规模模型的标准姿势。
它的局限,和它留下的遗产
8-bit Adam 当然不是万能药。它的实现比普通 Adam 复杂得多,量化带来的超参数敏感性需要额外调参;收益也跟模型规模强相关——模型越大,优化器状态占比越高,省得越多,小模型上省的那点显存可能不值当。
但真正重要的,是它对后续工作的影响。block-wise 量化的思想被 Dettmers 本人直接带进了 2023 年的 QLoRA——那个用 NF4 量化做大模型微调的方法,如今微调 70B 模型几乎默认跑在 4-bit / 8-bit 上。而 8-bit 优化器本身,也早已集成进 HuggingFace 生态,成了 bitsandbytes 库的标配。
进阶阅读
想深入了解的可以读原论文 8-bit Optimizers via Block-wise Quantization(ICLR 2022),以及 Dettmers 后来发表的 QLoRA 论文。前者讲透了量化误差的分析,后者展示了 block-wise 思想在参数高效微调中的威力。
原创文章,作者:guanweilu,如若转载,请注明出处:https://guanweilu.cn/article/332.html