分布式训练与显存:大模型为什么需要一整个集群
拆解训练时显存在装什么,解释为什么单卡放不下大模型,系统讲清数据并行、张量并行、流水线并行与 ZeRO 分片的切分维度,以及混合精度、梯度累积、激活检查点这些省显存的关键工程手段
- 弄清训练显存由哪几部分构成,理解优化器状态为何是大头
- 区分数据并行、张量并行、流水线并行各自切什么
- 理解 ZeRO 分片在数据并行框架内省显存的思路
- 掌握混合精度、梯度累积、激活检查点分别用什么换什么
显存到底被谁吃掉了
训练一个模型,显存远不止装参数本身。以常见混合精度训练为例,主要有四块:参数、梯度、优化器状态、前向保存的激活值。其中优化器状态常是最大头——例如 Adam 要为每个参数保存一阶矩、二阶矩,再叠加主权重,参数量乘以若干倍才是真实占用。
粗略估算,用 Adam 混合精度训练,每一个模型参数可能占用约 16~20 字节(参数+梯度+优化器状态,未计激活)。因此数十亿到数千亿参数的模型,单张消费级显卡根本放不下,必须把模型和计算「切」到多卡多机——这就是分布式训练。
三种并行:沿不同维度切
每卡放完整模型副本、喂不同数据分片
各卡算局部梯度,再 AllReduce 同步平均
最简单;但单卡要装得下整个模型
把单层的权重矩阵切到多卡
单算子跨卡协作,卡间通信频繁
适合单机内高速 NVLink,突破单层显存
按层把模型竖切到不同节点
数据分微批像流水线一样流动
适合跨机;要处理流水线气泡
实际大模型训练几乎总是 3D 并行:节点内用张量并行、节点间用流水线并行、整体再叠加数据并行,在显存、通信和负载之间找平衡。
ZeRO:在数据并行内部再分片
普通数据并行每张卡都冗余存一份优化器状态、梯度甚至参数,浪费严重。ZeRO 的思路是把这些状态在各卡间分片,需要时再通过通信取回,分阶段(ZeRO-1 切优化器状态、-2 再切梯度、-3 连参数也切)逐级省显存,用可控的通信换显存,是当前训练框架的主流基座。
省显存与稳训练的常用手段
混合精度
用 FP16/BF16 计算与存储、关键处保留 FP32 主权重,省显存且更快,BF16 数值更稳
梯度累积
多个小批的梯度累加后再更新一次,用小显存模拟大 batch
激活检查点
不保存全部中间激活,反向时按需重算,用额外计算换大显存节省
梯度裁剪
给梯度范数设上限,防止偶发尖峰导致训练发散
切分省了显存,但带来卡间通信:TP 最频繁(层内)、DP 需要梯度同步、PP 有流水线气泡。集群用高速互联(NVLink、InfiniBand)正是为此。设计并行方案的本质,是在显存容量、计算效率和通信开销三者间做工程权衡,没有免费午餐。
在 Adam 训练中,通常占用显存最大、且常被初学者忽略的部分是?
资深工程师加餐
底层原理 · 大厂视角 · 工程经验,点卡片展开
先查数据:标签是否正确、有没有特征和标签泄漏、是否做了归一化;再查损失与学习率:学习率太大震荡、太小几乎不动;接着看批量大小和优化器,最后才怀疑模型结构。一个极好用的冒烟测试是:用极少量样本训练,看损失能不能被压到接近零——能,说明整条代码链路是通的;不能,问题一定在实现里,先别谈调参。