跳转到内容

输入关键词开始搜索

    资料摘要:personal chatgpt 13 — gradient checkpointing 显存优化 trick

    视频摘要更新 2026-08-02置信度 high待阅读原始来源 ↗#深度#大模型#课程#personal-chatgpt

    本讲由真实视频转录与作者 notebook/代码交叉整理。原始层:🔒 带时间戳视频转录;课程总览:personal chatgpt — LLMs 实践系列

    反向传播需要前向阶段的中间激活。梯度检查点不保存全部激活,只保留部分检查点,并在反向传播时重算缺失的中间激活,以额外计算时间换取更低的激活显存。它与第 11 讲的梯度累积作用位置不同:前者改变激活保存策略,后者改变参数更新频率。

    课程从简单计算图出发,说明反向传播求梯度时需要前向过程中的中间节点。网络深度增加时,保存全部 intermediate activations 的显存随深度增长。

    Notebook 引用的说明把两端策略概括为:保存全部激活时约需 O(n)O(n) 激活内存;极端重算策略可接近 O(1)O(1) 激活内存,但计算步数可能增至 O(n2)O(n^2)。实际 checkpointing 位于二者之间。

    只缓存若干边界激活。反向传播到某段时,从最近检查点重新执行该段前向,恢复计算梯度所需的中间值。显存降低,但反向阶段增加额外前向计算。

    与第 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。

    training_args = TrainingArguments(
    per_device_train_batch_size=1,
    gradient_accumulation_steps=4,
    gradient_checkpointing=True,
    **default_args,
    )

    概念性的资源权衡:

    更少保存的激活更低显存+更多重算\text{更少保存的激活}\quad\Longrightarrow\quad \text{更低显存}+\text{更多重算}

    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 保存输出。
    • 承接第 11 讲: 两讲共同回答“显存不够时如何训练”,但优化机制不同。
    • 引出第 14 讲: Llama 2 微调通常同时结合量化、PEFT、梯度累积与 checkpointing 等策略。
    • 运行层面的组合: 梯度累积与 checkpointing 可以一起开,但对吞吐、随机性和有效 batch 的影响应分别评估。
    1. 梯度检查点为什么能降低显存,又为什么会增加训练时间?
    2. 它主要减少权重、optimizer state 还是激活的显存?
    3. 本讲实验为什么不是严格的单变量对照?
    4. 梯度检查点与梯度累积能否同时使用?它们分别改变什么?
    • 滴答清单抓取时学习状态:已完成
    • 视频正文采用自动语音识别;关键术语已用标题、notebook 与源码校正,无法确认的口语细节不扩写。
    • notebook 未执行的 CUDA、权重下载、训练或外部 API 单元,不表述为已复现实验。