为什么 70B 模型一张 GPU 放不下?
七十B参数,如果只看FP16权重,大约需要140GB,训练时远远不只是保存模型权重,还需要模型参数,梯度,optimizer state,activation,所以训练显存远远大于140GB.
训练比推理更吃显存,推理时:
权重 + Activation + KV Cache训练时:显存不仅被模型参数占用,还会被梯度、优化器状态和中间激活占用。
权重 + 梯度 + Optimizer State + Activation怎么解决呢?
把工作拆到多张GPU上,
怎么拆呢?
有这么多的并行方式:
Data Parallelism Model Parallelism Tensor Parallelism Pipeline Parallelism ZeRo FSDP我们现在逐一介绍:
DP:每张 GPU 都保存一份完整模型,然后不同 GPU 处理不同的数据,一张GPU处理Batch A,现在四张GPU就可以同时处理ABCD,所以可以训练吞吐量。但是这种方法有一个致命缺点:即使所有 GPU 加起来的总显存够,单张 GPU 放不下模型,DP 依然无法解决问题。
完整模型 │ ┌─────────┼─────────┐ ↓ ↓ ↓ GPU 0 GPU 1 GPU 2 ... │ │ │ Batch A Batch B Batch CZero Redundancy Optimizer(ZeRo):减少 Data Parallelism 中重复保存的数据,训练需Parameters,Gradients,Optimizer States,它大量重复,把这些状态切分到不同的GPU,这样减少冗余显存。DeepSpeed 它是微软开源的一个大模型训练与推理优化框架,它提供了很多分布式能力,最著名的就是zeRo
ZeRo-1:切分optimizer states
Parameters ↓ 每张 GPU 仍然有 Gradients ↓ 每张 GPU 仍然有 Optimizer States ↓ 分片保存ZeRo-2:再ZeRo-1的基础上,再切分Gradients.ZeRo-3同理,这里就不详细赘述了。则合理说一下Fully Shared Data Parellel,它是PyTorch 提供的一种参数、梯度和优化器状态分片的分布式训练方案,它的思想和zero-3很像
| 技术 | Parameters | Gradients | Optimizer States |
|---|---|---|---|
| DP | 复制 | 复制 | 复制 |
| ZeRO-1 | 复制 | 复制 | 分片 |
| ZeRO-2 | 复制 | 分片 | 分片 |
| ZeRO-3 | 分片 | 分片 | 分片 |
Model Parallelism:把一个模型本身拆到多张 GPU,但是GPU 之间需要频繁通信。
GPU 0 → 模型的一部分 GPU 1 → 模型的一部分 GPU 2 → 模型的一部分 GPU 3 → 模型的一部分PP:按照模型层,把模型拆成多个 Stage,然后通过Micro-batch让不同 Stage 像流水线一样工作。
GPU 0 → Layer 1~8 GPU 1 → Layer 9~16 GPU 2 → Layer 17~24 GPU 3 → Layer 25~32然后再把batch再拆成多个Micro-batch,GPU最大程度使用
Batch ↓ Micro-batch 1 Micro-batch 2 Micro-batch 3 Micro-batch 4Pipeline Bubble就是最后总是会出现GPU完成工作以后空闲的情况,所以 Micro-batch 越合理,通常越能降低 Bubble 带来的效率损失。
3D Parallelism:
我们已经有DP,TP,PP,我们可以同时使用,大型LLM训练通常就是多种并行方式组合使用,这就是3D Parallelism
Data Parallel × Tensor Parallel × Pipeline ParallelLLM Training │ ┌──────────┼──────────┐ ↓ ↓ ↓ DP TP PP │ │ │ 数据拆分 张量拆分 层拆分TP:把同一个 Layer 内部的矩阵计算拆到多张 GPU。
Y=XW,其中X是输入,W是一个非常大的权重矩阵,如果W太大,一张GPU放不下,可以把W切开,
eg:
GPU 0: X × W₁ GPU 1: X × W₂ GPU0 +GPU1 = 完整结果我们来总结一下:
LLM │ ┌────────┴────────┐ ↓ ↓ 推理 训练 │ │ ┌─────┼─────┐ ┌────┼─────┐ ↓ ↓ ↓ ↓ ↓ ↓ KV Flash Quant DP TP PP Cache Attention ization │ │ │ │ └────┼─────┘ ↓ ↓ PagedAttention ZeRO │ │ └──────→ vLLM ↓ DeepSpeed / FSDP