42 分钟
大模型内核:从经典机器学习到对齐与前沿

分布式训练与显存:大模型为什么需要一整个集群

拆解训练时显存在装什么,解释为什么单卡放不下大模型,系统讲清数据并行、张量并行、流水线并行与 ZeRO 分片的切分维度,以及混合精度、梯度累积、激活检查点这些省显存的关键工程手段

  • 弄清训练显存由哪几部分构成,理解优化器状态为何是大头
  • 区分数据并行、张量并行、流水线并行各自切什么
  • 理解 ZeRO 分片在数据并行框架内省显存的思路
  • 掌握混合精度、梯度累积、激活检查点分别用什么换什么

显存到底被谁吃掉了

训练一个模型,显存远不止装参数本身。以常见混合精度训练为例,主要有四块:参数、梯度、优化器状态、前向保存的激活值。其中优化器状态常是最大头——例如 Adam 要为每个参数保存一阶矩、二阶矩,再叠加主权重,参数量乘以若干倍才是真实占用。

训练显存的四大组成(自底向上叠加)
激活值前向为反向传播临时保存的中间结果,随 batch×序列长度×层数增长
优化器状态Adam 动量/二阶矩+主权重,通常是参数的数倍,显存大头
梯度与参数同形,每个参数一份
模型参数权重本身(混合精度下每参数 2 字节)
ℹ️一个数量级直觉

粗略估算,用 Adam 混合精度训练,每一个模型参数可能占用约 16~20 字节(参数+梯度+优化器状态,未计激活)。因此数十亿到数千亿参数的模型,单张消费级显卡根本放不下,必须把模型和计算「切」到多卡多机——这就是分布式训练。

三种并行:沿不同维度切

三大并行策略切的是什么
数据并行 DP

每卡放完整模型副本、喂不同数据分片

各卡算局部梯度,再 AllReduce 同步平均

最简单;但单卡要装得下整个模型

张量并行 TP

把单层的权重矩阵切到多卡

单算子跨卡协作,卡间通信频繁

适合单机内高速 NVLink,突破单层显存

流水线并行 PP

按层把模型竖切到不同节点

数据分微批像流水线一样流动

适合跨机;要处理流水线气泡

实际大模型训练几乎总是 3D 并行:节点内用张量并行、节点间用流水线并行、整体再叠加数据并行,在显存、通信和负载之间找平衡。

ZeRO:在数据并行内部再分片

普通数据并行每张卡都冗余存一份优化器状态、梯度甚至参数,浪费严重。ZeRO 的思路是把这些状态在各卡间分片,需要时再通过通信取回,分阶段(ZeRO-1 切优化器状态、-2 再切梯度、-3 连参数也切)逐级省显存,用可控的通信换显存,是当前训练框架的主流基座。

省显存与稳训练的常用手段

四个高频工程手段

混合精度

用 FP16/BF16 计算与存储、关键处保留 FP32 主权重,省显存且更快,BF16 数值更稳

梯度累积

多个小批的梯度累加后再更新一次,用小显存模拟大 batch

激活检查点

不保存全部中间激活,反向时按需重算,用额外计算换大显存节省

梯度裁剪

给梯度范数设上限,防止偶发尖峰导致训练发散

🐍通信是分布式的隐藏瓶颈

切分省了显存,但带来卡间通信:TP 最频繁(层内)、DP 需要梯度同步、PP 有流水线气泡。集群用高速互联(NVLink、InfiniBand)正是为此。设计并行方案的本质,是在显存容量、计算效率和通信开销三者间做工程权衡,没有免费午餐。

选择题

在 Adam 训练中,通常占用显存最大、且常被初学者忽略的部分是?

资深工程师加餐

底层原理 · 大厂视角 · 工程经验,点卡片展开

先查数据:标签是否正确、有没有特征和标签泄漏、是否做了归一化;再查损失与学习率:学习率太大震荡、太小几乎不动;接着看批量大小和优化器,最后才怀疑模型结构。一个极好用的冒烟测试是:用极少量样本训练,看损失能不能被压到接近零——能,说明整条代码链路是通的;不能,问题一定在实现里,先别谈调参。