在大模型训练中,"这批参数到底能不能塞进显卡"几乎是每个工程师上手第一件事就要面对的问题。很多时候一次 OOM(Out of Memory)背后隐藏的不是"显卡不够大",而是对显存构成理解不够精细——参数、梯度、优化器状态、激活值这四块各自占多少,混合精度和 ZeRO 又是怎么把这些数字重新分配到多张卡上的。本文从显存构成的第一性原理出发,逐步展开到 DeepSpeed 的工程实现,最后落到 DDP / device_map 等实战中最常踩的坑。一、显存构成:参数、梯度、优化器、激活值训练时显卡上的显存大体可以拆成四块:总显存 ≈ 模型参数 + 梯度 + 优化器状态 + 激活值 + 显存碎片/预留前三项是"静态"的,只跟模型结构和优化器设置有关,跟 batch size、seq_len序列长度无关。激活值则是"动态"的,随batch size、seq_len、层数线性甚至更高阶增长。这也是为什么很多人发现模型加载没问题,一开始训练就爆显存,这爆的往往是激活值。下面统一用Ψ表示模型参数量(例如 7B 模型 Ψ = 7×10⁹)。1.1 模型参数显存每个参数占用的字节数取决于存储精度:精度每参数字节数说明FP324 Bytes传统全精度训练FP16 / BF162 Bytes混合精度训练中的"计算精度"INT81 Byte量化推理场景INT40.5 Byte极限量化纯 FP32 训练时,参数显存 = 4Ψ。如果用混合精度训练(AMP),通常会同时保留一份 FP16/BF16 的权重(用于前向/反向计算)和一份 FP32 主权重(用于优化器更新),这部分会在下面的"优化器状态"里一起算,避免重复计数。混合精度训练的初衷是用 FP16/BF16 做前向和反向计算(省显存、算得快),但如果优化器更新也直接在 FP16/BF16 权重上进行,会出现一个数值问题:权重更新量经常小到被舍入吃掉。1.2 梯度显存梯度的精度一般和前向计算精度保持一致:纯 FP32 训练:梯度 = 4Ψ混合精度训练:梯度通常以 FP16/BF16 存储 = 2Ψ(部分实现如 DeepSpeed ZeRO 会额外维护一份 FP32 梯度用于累积,视配置而定)1.3 优化器状态显存以最常用的Adam / AdamW为例,优化器需要为每个参数维护一阶动量m和二阶动量v,这两者出于数值稳定性考虑几乎总是用 FP32 存储,哪怕主训练用的是混合精度。这就引出了 ZeRO 论文中那个经典的"16Ψ" 公式(混合精度 + Adam 场景):FP16 参数: 2Ψ FP16 梯度: 2Ψ FP32 主参数: 4Ψ FP32 一阶动量 m: 4Ψ FP32 二阶动量 v: 4Ψ ------------------------ 合计: 16Ψ也就是说,混合精度 + Adam训练一个 7B 模型,仅"参数+梯度+优化器状态"这三项静态显存就需要:16 × 7×10⁹ Bytes ≈ 112 GB这还没算激活值,已经远超单张 80GB A100/H100 的容量——这正是为什么 7B 起步的模型几乎不可能用单卡朴素训练,必须依赖分布式技术。不同优化器的系数不同,简单对比一下(假设混合精度 FP16/BF16 + FP32 主权重):优化器额外状态优化器状态系数总系数(参数+梯度+优化器)