先这样答
训练时的显存主要分为两部分:静态的模型状态和动态的激活值。模型状态包含参数、梯度和优化器状态;激活值则是前向传播产生的中间结果。估算时通常先算静态部分的基准占用,再根据具体的训练配置叠加动态部分,最后为混合精度下的损失缩放等操作留出余量。
以常规的Adam混合精度训练为例,假设模型参数量为Θ。参数本身用FP16格式存储占2Θ字节,梯度同样是FP16格式占2Θ字节。优化器状态是整个静态占用中最重的部分,Adam算法需要保存FP32格式的参数主副本、一阶动量和二阶动量,这三者分别占用4Θ、4Θ和4Θ字节,共计12Θ字节。综合下来,静态模型状态总共需要约16Θ字节。这意味着一个7B参数规模的模型,仅静态部分就需要上百GB显存,单卡通常是无法容纳的。
激活值的估算机制与静态部分不同。它与序列长度、批次大小、网络层数呈线性增长关系,并与隐藏层维度成正比。如果不使用Flash Attention机制,注意力计算的激活值还会随序列长度呈平方级增长。在长序列训练场景中,激活值占用的显存往往会超过模型参数与优化器状态的总和。这也是为什么在长序列训练中,必须采用重算机制来用计算时间换取显存空间。
在面试官面前可以这样收束:静态部分直接按16倍参数量估算,动态部分指出长序列场景下激活值才是占用核心。实际工程中应对显存压力通常采用两级方案,先使用梯度检查点换取激活显存,再配合ZeRO技术将模型状态分片到多张显卡上。
面试官会怎么追问
- 「遇到OOM报错,常用的显存优化技术有哪些?」 应对显存不足可以按代价从小到大的顺序采取措施。首先开启Flash Attention直接消除注意力激活的平方项,并减小批次大小配合梯度累积。其次开启梯度检查点,牺牲约三成的重算时间省下大量激活显存。接着引入ZeRO分片技术将模型状态切分到多卡,如果仍然不够,最后再考虑代价较大的CPU卸载或采用低精度优化器与量化训练。
- 「为什么长序列训练时激活值会成为显存瓶颈?」 因为模型状态的显存占用是固定的,而前向传播保存的中间激活值会随序列长度增加而膨胀。特别是标准注意力矩阵的计算包含序列长度的平方项,当上下文窗口拉长时,这部分中间变量的显存占用会迅速超过参数和优化器状态的总和。
- 「梯度检查点的具体机制是什么?」 它的核心机制是用计算换显存。前向传播时不保存所有层的激活值,而是每隔几层保存一个断点。反向传播执行到未保存激活值的层时,系统会从最近的断点重新执行前向计算来生成需要的中间结果。这能大幅降低峰值显存,代价是增加了额外的计算时间。
回答的坑
- 忽略优化器状态的精度差异,错误地把所有状态都按FP16计算,应该明确Adam优化器需要保留FP32的参数主副本和两个动量。
- 估算激活值时漏掉序列长度的影响,只盯着模型参数量看,正确做法是指出长序列场景下激活值才是显存占用的大头。
同系列的题
—— 本题完 ——