资料摘要:personal chatgpt 13 — gradient checkpointing 显存优化 trick
本讲由真实视频转录与作者 notebook/代码交叉整理。原始层:🔒 带时间戳视频转录;课程总览:personal chatgpt — LLMs 实践系列。
反向传播需要前向阶段的中间激活。梯度检查点不保存全部激活,只保留部分检查点,并在反向传播时重算缺失的中间激活,以额外计算时间换取更低的激活显存。它与第 11 讲的梯度累积作用位置不同:前者改变激活保存策略,后者改变参数更新频率。
1. 为什么激活占显存
Section titled “1. 为什么激活占显存”课程从简单计算图出发,说明反向传播求梯度时需要前向过程中的中间节点。网络深度增加时,保存全部 intermediate activations 的显存随深度增长。
Notebook 引用的说明把两端策略概括为:保存全部激活时约需 激活内存;极端重算策略可接近 激活内存,但计算步数可能增至 。实际 checkpointing 位于二者之间。
2. 检查点与重算
Section titled “2. 检查点与重算”只缓存若干边界激活。反向传播到某段时,从最近检查点重新执行该段前向,恢复计算梯度所需的中间值。显存降低,但反向阶段增加额外前向计算。
3. Transformers 实验
Section titled “3. Transformers 实验”与第 11 讲类似,实验使用 bert-large-uncased 与随机分类数据。无 checkpoint:约 23.78 秒、21.53 samples/s、12,917 MB;启用 checkpoint:约 38.96 秒、13.14 samples/s、9,305 MB。保存结果明确显示显存下降、速度变慢。
注意:checkpoint 版本同时把 micro-batch 从 4 改为 1,并设置 gradient_accumulation_steps=4。因此该对比没有只改变一个变量,不能把全部显存/速度差异都归因于 checkpointing。
公式 / 代码
Section titled “公式 / 代码”training_args = TrainingArguments( per_device_train_batch_size=1, gradient_accumulation_steps=4, gradient_checkpointing=True, **default_args,)概念性的资源权衡:
PyTorch 还提供 torch.utils.checkpoint.checkpoint_sequential,Notebook 仅导入它,没有构造完整的手写分段示例。
| 时间 | 内容 |
|---|---|
| 00:42 | 从简单计算图解释中间节点与反向传播 |
| 03:05 | 扩展到更深计算图与激活保存问题 |
| 06:27 | 引入 checkpoint 节点 |
| 09:55 | 讨论 intermediate activations 与重算 |
| 11:42 | 开始比较不开启与开启 checkpoint 的训练 |
| 12:47 | 口述开启后显存约 9GB |
| 13:35 左右 | 总结显存与计算时间的 trade-off |
- 运行标签:conditional。 需要 CUDA、
bert-large-uncased下载、PyTorch、Transformers、Datasets 和 pynvml。 - 会执行两次真实训练并写入
tmp。 - 对照实验同时更改 micro-batch 与梯度累积参数,属于教学演示,不是严格消融实验。
- 当前 Transformers 版本的
evaluation_strategy、checkpoint 行为和警告可能与 Notebook 保存环境不同;依赖未锁定。 - 本次未重新运行 GPU 训练,数值均来自 Notebook 保存输出。
与前后讲关系
Section titled “与前后讲关系”- 承接第 11 讲: 两讲共同回答“显存不够时如何训练”,但优化机制不同。
- 引出第 14 讲: Llama 2 微调通常同时结合量化、PEFT、梯度累积与 checkpointing 等策略。
- 运行层面的组合: 梯度累积与 checkpointing 可以一起开,但对吞吐、随机性和有效 batch 的影响应分别评估。
- 梯度检查点为什么能降低显存,又为什么会增加训练时间?
- 它主要减少权重、optimizer state 还是激活的显存?
- 本讲实验为什么不是严格的单变量对照?
- 梯度检查点与梯度累积能否同时使用?它们分别改变什么?
- 视频:gradient checkpointing 显存优化 trick
- 转录:🔒 personal chatgpt 13 视频转录
- 课件 / 代码:🔒 gradient_checkpointing.ipynb
- 滴答清单抓取时学习状态:已完成。
- 视频正文采用自动语音识别;关键术语已用标题、notebook 与源码校正,无法确认的口语细节不扩写。
- notebook 未执行的 CUDA、权重下载、训练或外部 API 单元,不表述为已复现实验。