↓ Skip to main content
  1. Posts/

从推理到训练:模型调度中的显存约束

Note: This article is available in Chinese only. 本文暂无英文版本。 View original

推理任务能放下,为什么训练任务放不下
#

上一篇讨论了模型调度中的 Capacity Constraint:一组模型到底能不能同时放进一张 GPU。

那篇的结论是:推理场景下,判断模型能否部署到 GPU,先看 persistent memory(主要是模型权重),再加上 forward 过程中的 activation、workspace 和运行时开销。最终关心的是某一时刻"同时还活着"的 tensor 有多少——也就是 peak live memory。

但如果把其中一个推理任务换成训练任务,这套估算很快就会失效。

同一个模型,参数没有变,推理时可能只需要几个 GB。一旦开始训练,显存可能膨胀到推理的数倍甚至十几倍。更麻烦的是,训练显存甚至不是一个稳定的数字——它在 forward、backward 和 optimizer step 之间不断变化,峰值出现在哪个阶段都不一定。

这里有一个看起来反直觉的现象:activation checkpointing 会主动丢掉 forward 已经算出来的中间结果,在 backward 时重新计算一遍。“把计算做两遍"怎么反而是一种优化?

理解这个问题,需要先理解训练显存到底由什么决定。

训练显存账本多了什么
#

上一篇建立了推理场景的 心智模型:

\[\boxed{M_{\text{infer, peak}} \approx M_{\text{persistent}} + M_{\text{peak\ runtime\ working\ set}}}\]

训练在这个基础上多出了两块开销。

第一块是模型状态变大了。推理只需要保存模型权重。训练还需要:每个参数对应一份 gradient;以及 optimizer 的内部状态(比如 Adam 需要维护一阶矩和二阶矩,都是 FP32)。

粗估一下数量级。BF16 推理主要是约 2 bytes/parameter 的模型权重。而 BF16 mixed-precision Adam 训练,仅模型状态(参数 + 梯度 + optimizer states)就可能接近 16 bytes/parameter——同一个模型,参数量没变,训练的模型状态已经是推理的 8 倍。

第二块更关键:为了 backward 而保留的 forward 中间结果。这才是训练显存暴涨的主要推手。

Backward 为什么要留住 forward 的中间结果
#

先用最简单的线性层看清楚这件事。

Forward:

\[Y = WX\]

Backward 中,已知上游传回的 \(\frac{\partial L}{\partial Y}\),需要计算两个梯度:

\[\frac{\partial L}{\partial W} = \frac{\partial L}{\partial Y} X^T\]\[\frac{\partial L}{\partial X} = W^T \frac{\partial L}{\partial Y}\]

注意第一个式子:计算 parameter gradient \(\frac{\partial L}{\partial W}\) 需要 forward 阶段的输入 \(X\)。

这意味着 \(X\) 不能在 forward 结束后立即释放。它必须一直活到对应的 backward pass 完成。

这就是训练显存和推理显存最根本的区别。推理时,一个 activation 被下一层使用后通常可以立即释放。训练时,backward 强制延长了 activation 的 lifetime。

而且不只是一层。一个 N 层网络的第一层输入,可能要等到 backward 一路回传到第一层时才能释放——也就是整个 forward 加上几乎整个 backward 期间都得常驻。

需要明确一个术语:训练系统语境中的 activation,不只是激活函数的输出,而是所有为了 backward 而保留的前向中间结果。 Autograd 框架会根据 computation graph 的依赖关系,自动决定哪些中间 tensor 需要保存。

显存是一条时间线
#

上面的分析说明了训练比推理多出哪些状态。但仅仅知道这些状态各有多大还不够——更关键的问题是:它们什么时候同时存在?

一次训练 iteration 中,显存占用是随执行阶段变化的:

  • Parameters 和 optimizer states 几乎全程常驻。
  • Forward 阶段:activations 不断积累,越来越多的中间结果被保留等待 backward。
  • Backward 阶段:gradients 逐层产生,同时对应层的 saved activations 被使用后逐渐释放。
  • Optimizer step:parameters、gradients 和 optimizer states 同时参与更新,可能产生新的临时峰值。
训练 Iteration 显存生命周期与依赖

观察这条时间线,有几件事变得很清楚。

Forward 和 backward 交界处附近,大量 activations 和刚开始产生的 gradients 同时存在——这通常是整个 iteration 中显存占用最高的区域。

OOM 看的正是时间轴上的峰值,不是平均显存,也不是模型加载后的静态占用。

Activation Checkpointing 改变了什么
#

现在可以回到开头那个反直觉的问题:为什么"把计算做两遍"反而是一种优化?

上面的时间线图已经展示了问题所在:早期层的 activation(如 Act_L1)从 forward 一开始就产生,一直要等到 backward 几乎结束才能释放。存活时间非常长。

Activation checkpointing 的做法是:

  • 只保留少数边界 activation(checkpoint 点)。
  • 丢弃 checkpoint 区间内部的中间结果。
  • Backward 到达该区间时,从边界 activation 出发,重新执行部分 forward。
  • 使用重新生成的中间结果完成 backward,随后立即释放。

仅仅说"以计算换显存"还不够精确。Activation checkpointing 没有改变模型参数——参数量不变、单个 tensor 不变。它改变的是:

  • 哪些 tensor 需要长期保存。
  • Tensor 何时重新生成。
  • 同一时刻有多少 tensor 同时存活。

Activation checkpointing 真正缩短的是中间 tensor 的存活时间。

从 FLOPs 角度看(FLOPs 那篇讨论过 forward 约 \(F\)、backward 约 \(2F\)、正常训练约 \(3F\)):极端情况下完整重算一次 forward,总计算量约为 \(4F\),增加约三分之一。实际开销取决于 checkpoint 粒度和 kernel 效率。

同时需要认清 checkpointing 的边界:

  • 主要减少 retained activations。
  • 不减少 parameters 和 optimizer states。
  • 不一定降低单个算子执行时的 workspace 峰值。

因此,如果仅模型状态(parameters + optimizer states)就已经放不下一张 GPU,activation checkpointing 无法解决根本问题。

从显存曲线回到模型调度
#

推理场景下,上一篇已经指出 model.memory 是一个危险的抽象。到了训练场景,这个问题更加严重。

一个调度器判断训练任务能否放到某张 GPU 上,不能只记录 \(\text{model\_size} = X \text{ GB}\)。它至少还需要理解:

  • 常驻模型状态的大小(parameters + gradients + optimizer states)。
  • 当前 micro-batch 和序列长度下,activation 的峰值有多大。
  • Activation checkpointing 是否开启——同一个模型,开不开 checkpointing,显存曲线的形状完全不同。
  • Forward、backward、optimizer step 各阶段的峰值分别出现在哪。

训练任务的资源需求不是一个静态的"模型大小”,而是一条由执行计划决定的显存曲线。

同样的模型、同样的 GPU,开了 checkpointing 可能放得下,不开就 OOM——但模型参数量没有任何变化。

总结
#

上一篇建立的推理 心智模型 是 Persistent Memory + Peak Runtime Working Set。到了训练阶段,这种静态视角就不够了。

训练相对推理的核心差异:backward 强制延长了 forward 中间结果的存活时间。 显存峰值出现在 forward-backward 交界处,大量 activations 和刚开始产生的 gradients 同时存在。OOM 看的是这条时间线上的峰值。

Activation checkpointing 的本质是缩短中间 tensor 的存活时间。 它没有减少参数或 optimizer states——它改变的是哪些 tensor 需要同时存在。

模型大小只是模型本身的属性;训练显存是执行计划的属性。 调度器不能用一个静态的 model_size 来判断训练任务能否放下。

参考资料
#

Related