跳到正文

显存优化技术原理与对比

梯度检查点、模型并行、量化等关键技术详解

原题:在大模型训练或推理过程中,显存不足是一个常见瓶颈。请列举并解释有效的显存优化技术,如梯度检查点、模型并行、量化等。

推理优化 · 海尔真题

回答与解析

第一步不是选技术,而是做内存账本

训练显存至少拆成 parameters、gradients、optimizer states、activations、临时 workspace/通信 buffer;推理拆成 weights、KV Cache、activations/workspace 与调度开销。谁占主导取决于参数量、优化器、序列长度、batch、精度和并行方式,不能笼统说激活永远最大。

训练侧

  • BF16/FP16 降低权重与计算精度;FP16 常配 loss scaling,BF16 因指数范围较大通常不需要同类缩放;
  • activation checkpointing 以反向重算换保存激活,收益和额外计算由 checkpoint 粒度及图决定;
  • ZeRO/FSDP 分片状态,tensor/pipeline/sequence parallel 拆参数或序列,选择取决于内存账本和通信;
  • FlashAttention 通过 IO-aware tiling 避免物化完整注意力矩阵;
  • CPU/NVMe offload 以传输换显存,梯度累积以更多 micro-step 换较小 micro-batch。

推理侧

权重量化、KV Cache 量化/分页/卸载、GQA/MQA/MLA 架构,以及限制最大上下文和并发,可降低或约束显存峰值。continuous batching 主要改善动态调度和设备利用率,本身不保证降低峰值;更多请求的 KV 同时驻留时,峰值还可能上升,需配合 paged KV、最大并发、batch token 或 memory budget 控制内存。无论采用哪种方案,质量、TTFT、TPOT 与吞吐都必须一起测。

不要给 checkpoint 固定节省百分比、INT4 固定精度损失或固定提速,也不要按“8 卡以下 ZeRO-2、更多 ZeRO-3”选型。先用 profiler 与公式估算,再在目标硬件上测峰值、OOM margin 和性能。

口语版讲法(约4分钟)

  • 先拆训练与推理内存账本
  • 按主导项选择训练侧技术
  • 解释 checkpoint 与分片的代价
  • 说明推理侧权重与 KV 优化
  • 以目标硬件 benchmark 验收

显存优化的第一步不是背技术清单,而是建立内存账本。训练时我会拆成参数、梯度、优化器状态、激活、临时算子 workspace 和通信 buffer;推理时拆成模型权重、KV Cache、当前激活或 workspace 以及调度开销。谁最大没有固定答案:短序列、全参数 Adam 训练可能是参数和优化器状态主导;长序列、大 micro-batch 时激活可能主导;长上下文高并发推理则常由 KV Cache 主导。

如果训练侧是精度和状态占用高,先看 BF16 或 FP16。FP16 指数范围小,常需要动态 loss scaling;BF16 指数范围更大,通常不需要同类缩放,但是否溢出仍要监控。若激活占用高,可以用 activation checkpointing,只保存部分边界,反向时重算中间结果。它是用额外计算换显存,具体节省和时间增量取决于网络、序列长度和 checkpoint 粒度,不能写固定百分比。

若参数、梯度和优化器状态放不下,可以用 ZeRO 或 FSDP 分片;单层本身放不下时考虑 tensor parallel,层很多时可结合 pipeline parallel,长序列还可能用 sequence parallel。阶段和并行度不是按“几张卡”机械选择,而要看被分片对象、网络拓扑、通信占比和峰值 all-gather。CPU 或 NVMe offload 能继续省 GPU 显存,但会受 PCIe、主存或存储吞吐限制。

注意力部分,FlashAttention 通过 IO-aware tiling,不物化完整注意力矩阵,能降低中间内存并提高算子效率。梯度累积把大 batch 拆成多个 micro-step,降低单步激活峰值,但总计算不会凭空消失。

推理侧先看权重与 KV 各占多少。权重量化可以降常驻内存,KV 可以做分页管理、量化或卸载;GQA、MQA、MLA 是架构级降低 KV 的方式。Continuous batching 主要提高动态调度和设备利用率,本身不保证降低峰值;更多并发请求的 KV 同时驻留时,峰值还可能上升,需要配合 paged KV、最大并发和 batch token 或 memory budget 约束。它也可能改变单请求排队和延迟。INT4 的精度损失、速度收益和 kernel 支持都依模型、任务和硬件,不能承诺固定小于百分之一或提速几倍。

验收时我会记录模型、dtype、batch、输入输出长度和并行配置,用 profiler 看各组成项,测峰值显存、OOM 安全余量、step time 或 TTFT/TPOT、吞吐和质量回归。优化顺序是先找主导项,再选择最小代价的技术,而不是把所有开关一起打开。

我也会逐项上线而不是一次叠满优化:先记录基线,再开 checkpoint、分片、量化或 KV 优化中的一项,分别测峰值、速度和质量。多种技术同时启用容易出现通信、重算和量化 kernel 相互抵消,最终显存省了却吞吐更差。

最后保留足够 OOM 安全余量,因为编译缓存、通信 bucket 和真实流量峰值可能不出现在单次离线测量里。线上峰值应覆盖最长请求和最高并发组合。

关键一句:显存主导项随参数、优化器、序列长度和并发变化,必须先做训练/推理分开的内存账本。

核验来源

  1. Training Deep Nets with Sublinear Memory Cost
  2. ZeRO: Memory Optimizations Toward Training Trillion Parameter Models
  3. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
  4. PyTorch Automatic Mixed Precision official documentation

面试官还可能这样问

  1. 问法 1 · 场景切入

    假设你正在训练一个70B的模型,单卡显存只有80G,跑起来直接OOM。你通常怎么一步步节省显存?从最简单的开始,比如调batch size、开混合精度,再到梯度检查点、ZeRO这些,说说你的优化顺序和理由。

  2. 问法 2 · 层层追问

    大模型训练时显存不够你一般怎么处理?……具体说说梯度检查点是怎么节省显存的,它的代价是什么?……那如果模型太大单卡放不下呢,你会考虑模型并行还是ZeRO?两者怎么选?……推理阶段显存瓶颈又不一样,比如长上下文时KV Cache会暴涨,你有什么优化手段?

  3. 问法 3 · 直球架构

    列举并解释大模型训练和推理中常见的显存优化技术,包括梯度检查点、模型并行、量化、ZeRO等。说清楚每种技术的原理、节省显存的效果以及各自的适用场景和trade-off。

同模块相关题目