ai-infra-interview-305

第 13 题:PyTorch的checkpoint机制(梯度检查点)如何节省显存?计算开销在哪里?

题目

PyTorch的checkpoint机制(梯度检查点)如何节省显存?计算开销在哪里?


完整讲解

一、为什么需要 checkpoint?

前向时,为了 backward 能算梯度,通常要把中间激活都存下来,显存占用 ≈ 与层数/序列长度成正比。大模型或长序列时,激活显存很容易成为瓶颈。梯度检查点(gradient checkpointing) 的思路是:前向时只存一部分激活(如每隔几层存一次),其余不存;反向时用到某段中间激活时,从最近的一个检查点重新算一遍那段前向,得到激活再算梯度。这样用重复计算显存


二、显存怎么省下来的?


三、计算开销在哪里?


四、PyTorch 里怎么用?


五、面试可怎么说?


面试要点


记忆要点

  1. 省显存:只存部分激活(检查点),其余反向时重算;显存从 O(L) 到 O(检查点数)。
  2. 开销:反向时多算分段前向,总计算量增加;trade-off 算力换显存。
  3. 使用:checkpoint(fn,…)、checkpoint_sequential 或按 block 包装。
返回模块 返回总览