跑得动大模型之前,先算清 AI 的显存账

2026-07-17 34 预计阅读时间: 1 分钟
来源: ruanyifeng.com AI 摘要 Original link

Disclaimer: This article is an AI-assisted summary. Read it together with the original source when precision matters. The summary may omit context, version differences, or edge cases and is not official documentation.

预计阅读时间:9 分钟

部署或训练 AI 模型时,参数量只是账单上的第一项。权重、KV Cache、中间激活、梯度、优化器状态和运行时工作区会在不同阶段争夺 GPU 显存。理解这些组成,才能回答三个实际问题:模型能否加载、请求能否达到目标并发、训练任务为什么突然 OOM。

由于来源摘要没有限定具体模型或框架,下面采用通用估算方法,并把代码明确作为可调整的实践示例。实际数字仍应以目标模型配置、推理引擎和硬件测量结果为准。

显存不是只有模型权重

最容易估算的是权重占用:

权重显存 ≈ 参数量 × 每个参数的字节数

常见存储格式可粗略按以下数值计算:

格式 理论字节/参数 70 亿参数的理论权重体积
FP32 4 约 26.1 GiB
FP16 / BF16 2 约 13.0 GiB
INT8 1 约 6.5 GiB
INT4 0.5 约 3.3 GiB

这些只是理论下限。量化模型还可能保存缩放因子、零点和分组元数据;推理框架也要分配 CUDA 上下文、临时张量和算子工作区。因此,一张 16 GiB 显卡并不意味着一定能稳定运行理论权重为 13 GiB 的模型。

推理阶段还要关注 KV Cache。自回归模型会缓存历史 token 的 Key 和 Value,避免每生成一个 token 都重新计算完整上下文。一个常用的粗略公式是:

KV Cache ≈ 层数 × 2 × token 数 × KV 头数 × 每头维度 × 每元素字节数 × 批大小

其中 2 代表 Key 和 Value。采用 GQA 或 MQA 的模型,其 KV 头数可能小于查询头数,所以不能只看注意力头总数。上下文长度和并发批大小都会近似线性推高 KV Cache。

训练为何比推理更“吃内存”

推理通常主要保留权重和 KV Cache;训练还要维护中间激活、梯度及优化器状态。以 Adam 类优化器为例,优化器通常需要保存一阶和二阶矩估计;某些混合精度方案还会维护 FP32 主权重。

因此,不能用“参数量乘以两字节”推断训练需求。训练显存更接近下面这张动态账单:

训练显存 = 权重 + 梯度 + 优化器状态 + 激活 + 临时工作区

激活占用会受批大小、序列长度、隐藏维度和网络层数影响。工程上常用以下手段交换显存、吞吐与实现复杂度:

  • 混合精度减少部分张量的存储和计算成本。
  • 梯度检查点通过重新计算部分前向结果来减少激活占用。
  • 梯度累积用多个小批次模拟更大的有效批次,但不会提升单步吞吐。
  • 张量并行、流水线并行或参数分片把状态分散到多张设备。
  • CPU 或 NVMe offload 能腾出 GPU 显存,但会引入传输延迟和带宽瓶颈。

动手估算权重与 KV Cache

可以这样实践:使用下面这个不依赖第三方库的 Python 脚本,在下载模型之前做第一轮容量评估。请把模型参数、层数、KV 头数和每头维度替换成目标模型配置文件中的值。

#!/usr/bin/env python3
import argparse

GIB = 1024 ** 3


def gib(num_bytes: float) -> float:
    return num_bytes / GIB


def main() -> None:
    parser = argparse.ArgumentParser(description="Estimate LLM weight and KV-cache memory")
    parser.add_argument("--params-b", type=float, required=True, help="Parameter count in billions")
    parser.add_argument("--weight-bytes", type=float, default=2.0, help="Bytes per weight")
    parser.add_argument("--layers", type=int, required=True)
    parser.add_argument("--kv-heads", type=int, required=True)
    parser.add_argument("--head-dim", type=int, required=True)
    parser.add_argument("--tokens", type=int, default=4096)
    parser.add_argument("--batch", type=int, default=1)
    parser.add_argument("--cache-bytes", type=float, default=2.0, help="Bytes per KV element")
    parser.add_argument("--overhead", type=float, default=1.15, help="Planning multiplier")
    args = parser.parse_args()

    weight_memory = args.params_b * 1_000_000_000 * args.weight_bytes
    kv_memory = (
        args.layers
        * 2
        * args.tokens
        * args.kv_heads
        * args.head_dim
        * args.cache_bytes
        * args.batch
    )
    subtotal = weight_memory + kv_memory

    print(f"Weights:        {gib(weight_memory):8.2f} GiB")
    print(f"KV cache:       {gib(kv_memory):8.2f} GiB")
    print(f"Subtotal:       {gib(subtotal):8.2f} GiB")
    print(f"Planned total:  {gib(subtotal * args.overhead):8.2f} GiB")
    print("Note: activations, CUDA context, allocator fragmentation, and workspaces may add more.")


if __name__ == "__main__":
    main()

例如,下面的命令估算一个假设模型:70 亿参数、32 层、8 个 KV 头、每头维度 128,使用 FP16 权重和 KV Cache,上下文长度 8192,批大小为 4。

python3 memory_estimator.py \
  --params-b 7 \
  --weight-bytes 2 \
  --layers 32 \
  --kv-heads 8 \
  --head-dim 128 \
  --tokens 8192 \
  --batch 4 \
  --cache-bytes 2 \
  --overhead 1.20

这个结果适合做容量筛选,不应视为框架承诺。动态批处理、分页式 KV Cache、张量并行和不同量化布局都会改变实际占用。

运行服务后,还应直接观察设备数据:

nvidia-smi --query-compute-apps=pid,process_name,used_memory \
  --format=csv -l 1

压测时同时记录输入长度、输出长度、并发数、吞吐和峰值显存。只有显存数字而没有负载条件,几乎无法用于下一次容量规划。

从 OOM 现象反推问题

模型启动时就 OOM,通常应先检查权重格式、设备映射以及框架初始化开销。短请求正常、长上下文失败,更可能与 KV Cache 或临时注意力张量有关。低并发正常、高并发失败,则要检查每个序列的缓存分配和批处理上限。

显存碎片也会制造“明明还有空间却无法分配”的现象。此时不要只看平均占用,还要查看峰值、最大连续分配需求和框架的内存统计。减少并发或上下文只是止血措施;生产环境还应设置输入 token 上限、输出 token 上限、队列长度和过载拒绝策略。

上线前的显存检查表

  1. 从模型配置中确认参数量、层数、KV 头数、头维度和数据类型。
  2. 分别估算权重常驻内存与单请求 KV Cache,不把二者混成一个数字。
  3. 为 CUDA 上下文、算子工作区、量化元数据和碎片预留余量。
  4. 使用接近生产分布的长短请求和并发组合进行压测。
  5. 记录峰值而非只记录空闲时或单请求的显存。
  6. 在量化、缩短上下文、降低并发和增加设备之间比较精度、延迟与成本。

AI 内存优化的核心并不是把模型勉强塞进显卡,而是让峰值占用可预测,让负载超过预算时能够受控退化。先建立一张能解释的显存账,再决定量化、并行或扩容,通常比遇到 OOM 后逐项试参数更可靠。


相关推荐