分布式深度学习


文档摘要

分布式深度学习 分布式训练把计算分散到多个 GPU 和多台机器上,用来训练那些对单设备而言太大或太慢的模型。本文件涵盖混合精度、数据并行、模型并行、流水线并行、ZeRO、FSDP、张量并行,以及 all-reduce 等通信原语——这些是大规模训练 LLM 的必备技术。 在单个 GPU 上训练一个大神经网络,终究会撞到墙:模型可能放不进内存,或者训练要花几个月。分布式训练把工作分散到多个设备(GPU、TPU 或整台机器)上,以更快地训练、训练更大的模型。本文件介绍让这一切成为可能的技术。 要理解为什么分布式重要,先从训练的计算代价说起。

分布式深度学习

分布式训练把计算分散到多个 GPU 和多台机器上,用来训练那些对单设备而言太大或太慢的模型。本文件涵盖混合精度、数据并行、模型并行、流水线并行、ZeRO、FSDP、张量并行,以及 all-reduce 等通信原语——这些是大规模训练 LLM 的必备技术。

  • 在单个 GPU 上训练一个大神经网络,终究会撞到墙:模型可能放不进内存,或者训练要花几个月。分布式训练把工作分散到多个设备(GPU、TPU 或整台机器)上,以更快地训练、训练更大的模型。本文件介绍让这一切成为可能的技术。

  • 要理解为什么分布式重要,先从训练的计算代价说起。一个有 d_{\text{in}} 个输入和 d_{\text{out}} 个输出的稠密层,在一批 B 个样本上做一次前向传播,大约需要 2 \cdot B \cdot d_{\text{in}} \cdot d_{\text{out}} FLOPs(浮点运算):输出矩阵每个元素一次乘法和一次加法。反向传播的代价大约是前向的两倍(要计算对输入和对权重的梯度),所以稠密层上一次训练步大约是 6 \cdot B \cdot d_{\text{in}} \cdot d_{\text{out}} FLOPs。

  • 对于一个隐藏维度为 d 的 Transformer 层,自注意力块涉及四个投影(Q、K、V 和输出),每个代价 O(B \cdot n \cdot d^2) FLOPs(n 是序列长度),再加上注意力矩阵计算的 O(B \cdot n^2 \cdot d)。前馈块有两个稠密层,通常先扩到 4d 再缩回来:O(B \cdot n \cdot 8d^2)。每层总计:大约 O(B \cdot n \cdot 12d^2 + B \cdot n^2 \cdot d)。乘上层数,你就能明白为什么训练 GPT 规模的模型需要成千上万的 GPU 小时。

  • 内存墙往往是更紧的约束。训练时,GPU 内存必须同时容纳四样东西:

堆叠柱状图,显示训练内存的构成:参数、梯度、优化器状态、激活

  • 参数:模型权重。一个 70 亿参数的模型用 FP32(每个参数 4 字节)光权重就需要 28 GB。

  • 梯度:与参数同等大小。又是 28 GB。

  • 优化器状态:Adam 维护两个额外的缓冲(一阶矩和二阶矩估计),每个都和参数一样大。为了数值稳定性,即使模型用低精度,它们也以 FP32 保存。对我们的 7B 模型,这就是 2 \times 28 = 56 GB。

  • 激活:前向传播时保存下来供反向传播使用的中间值。大小取决于批量大小、序列长度和模型宽度。这往往是最主要的成分,且随批量大小线性增长。

  • 我们这个用 FP32 Adam 的 7B 模型:28(参数)+ 28(梯度)+ 56(优化器)= 112 GB,还没算激活。单块 80 GB 的 A100 GPU 装不下。这就是为什么分布式策略必不可少。

  • **混合精度训练(mixed precision training)**是第一道防线。不再把所有东西都存成 FP32(32 位浮点),而是前向和反向传播用 FP16 或 BF16(16 位),同时为优化器更新保留一份 FP32 的主权重副本。

  • FP16 精度高(10 位尾数)但范围有限,可能上溢/下溢。损失缩放(在前向传播前把损失乘一个大因子,反向传播后再把梯度除以同样的因子)能缓解这一点。

  • BF16(brain float)与 FP32 有相同的指数范围(8 位指数),但精度更低(7 位尾数)。它几乎不溢出,也极少需要损失缩放,用起来更简单。BF16 是现代 Transformer 训练的默认选择。

  • 混合精度大约把激活和梯度的内存减半(它们是前向/反向传播时的大头),同时为数值稳定性把优化器状态留在 FP32。

  • **数据并行(data parallelism)**是最简单的分布式策略。你把整个模型复制到 N 个 GPU 上,把每个小批量均分成 N 份,每份发给一个 GPU。每个 GPU 独立地在自己那份上跑前向和反向传播。然后梯度在所有 GPU 间求平均(用 all-reduce 操作),每个 GPU 再更新自己本地的模型副本。

  • 从模型的视角看,这等价于用一个大了 N 倍的小批量来训练。如果每个 GPU 处理大小为 B 的批量,有效批量大小就是 N \cdot B

并排对比:数据并行复制模型、切分数据;模型并行切分模型、共享数据

  • 梯度求平均可以同步或异步地进行。同步 SGD 等所有 GPU 都完成再求平均,保证与单 GPU 上用更大批量训练在数学上等价。缺点是最慢的 GPU("掉队者")会拖住所有人。

  • 异步 SGD 让每个 GPU 独立地更新一个共享的参数服务器,无需等待。这消除了掉队者问题,但引入了"陈旧梯度":某个 GPU 可能在基于略微过时的参数上算梯度。陈旧梯度增加噪声,可能拖慢收敛。实践中,更倾向于用高效通信的同步 SGD。

  • **梯度累积(gradient accumulation)**是在有限硬件上模拟更大批量大小的一个软件技巧。你不再每个小批量做一次更新,而是跑若干次前向/反向传播并把梯度累积起来,最后做一次更新。这给出了与更大批量相同的结果,却不需要为激活准备更多 GPU 内存(同一时刻内存里只有一个小批量的激活)。

  • 当模型本身就大到放不进单个 GPU 时,你需要模型并行(model parallelism)。它主要有两种形式。

  • **张量并行(tensor parallelism)**把单个层切分到多个 GPU 上。一个大矩阵乘 Y = XW 可以按列切分:把 W 分成 [W_1, W_2] 分到两个 GPU 上,并行地算 Y_1 = XW_1Y_2 = XW_2,再拼接。这对注意力投影和前馈层都适用。它要求 GPU 之间通信极快(通常是节点内通过 NVLink),因为每一层都要合并部分结果。

  • **流水线并行(pipeline parallelism)**把不同的层分给不同的 GPU。GPU 0 跑第 1-4 层,GPU 1 跑第 5-8 层,依此类推。数据像流水线一样穿过。朴素的办法会有"流水线气泡":当 GPU 0 在处理微批次 1 的前向传播时,GPU 1-3 都闲着。**微批次(micro-batching)**通过把小批量切成更小的微批次依次穿过流水线来缓解这一问题,让所有 GPU 大部分时间都忙起来。

  • **混合并行(hybrid parallelism)**把数据、张量和流水线并行结合起来。一个典型的大模型设定可能在一个节点内(8 个由快速 NVLink 相连的 GPU)用张量并行、跨节点用流水线并行、跨节点组用数据并行。GPT-4 和 Llama 这类模型就是这么训练的。

  • 分布式训练的效率很大程度上取决于通信。关键操作是 all-reduce:给定 N 个 GPU 上各一个值,算出它们的和(或平均)并把结果分发给所有 GPU。

  • 朴素的 all-reduce 把所有数据发给一个 GPU,求和,再广播回去。这在通信上是 O(N),并在根节点形成瓶颈。

  • **环形 all-reduce(ring all-reduce)**高效得多。把 N 个 GPU 排成一个环。每个 GPU 把自己的数据切成 N 块。在 N - 1 步里,每个 GPU 把一块发给邻居、从另一边的邻居接收一块,累加部分和。再经过 N - 1 步,完整的和就传播到了所有 GPU。每个 GPU 传输的总数据量是数据大小的 2(N-1)/N 倍,当 N 增大时趋向 2\times。关键是,这并不随 N 增长,因而是带宽最优的。

四个 GPU 排成环,每个把梯度块传给邻居,直到所有 GPU 都拿到完整的和

  • **参数服务器(parameter server)**是另一种架构:专门的服务器节点持有模型参数。工作节点算出梯度发给服务器,服务器更新参数后再发回去。这更简单,但可能在服务器处形成通信瓶颈。

  • NCCL(NVIDIA Collective Communications Library,NVIDIA 集合通信库)是 GPU 间通信的标准库。它为 all-reduce、all-gather、broadcast 等集合操作提供优化实现,并自动为网络拓扑选择最佳算法。

  • **缩放定律(scaling laws)**描述模型性能如何随算力、数据和模型规模而提升。最初的 Kaplan 等人(2020)缩放定律发现,损失随这三者按幂律下降:

L(N) \propto N^{-\alpha_N}, \quad L(D) \propto D^{-\alpha_D}, \quad L(C) \propto C^{-\alpha_C}
  • 其中 N 是参数数量,D 是数据集大小,C 是算力预算。

  • Chinchilla 缩放定律(Hoffmann 等人,2022)表明,大多数模型都训练不足:对于给定的算力预算,你应该用一个更小的模型在更多的数据上训练,超出此前的认知。最优比例大约是每个参数 20 个词元。一个 7B 模型应该看到约 140B 个词元,而不是 Llama 1 用 65B 模型时用的 300B 词元。这一发现把整个领域推向了"算力最优"训练。

  • 混合专家(Mixture of Experts, MoE)是一种在不按比例增加计算的情况下扩大模型容量的架构。每个 Transformer 层不再只有一个前馈网络,而是有 N 个"专家"网络(每个都是一个标准 FFN)。一个门控网络(gating network)(路由器)审视每个词元,把它送到排名前 K 的专家(通常 K = 1K = 2)。

词元通过门控网络被路由到选中的专家,采用 top-K 稀疏路由并对输出加权组合

  • 总参数量大得多(因为有 N 个专家),但每个词元的 FLOPs 大致不变(因为每个词元只激活 K 个专家)。例如,Mixtral 8x7B 共有 47B 参数,但每次前向传播只动用约 13B,用更小模型的代价换来了大得多的模型的表现。

  • MoE 带来了挑战。负载均衡:如果路由器把大多数词元都送给同一个专家,其余专家就浪费了。一个辅助损失鼓励路由均匀。通信:不同专家可能住在不同 GPU 上,所以路由词元需要 all-to-all 通信,代价不菲。

  • 当训练持续数周或数月、跑在数千个 GPU 上时,**容错(fault tolerance)**至关重要。如果一块 GPU 挂了,你不想丢掉所有进度。**检查点(checkpointing)**定期把模型权重、优化器状态和训练状态(学习率、步数、数据位置)保存到磁盘。一旦发生故障,你就从最近的检查点重启。

  • 梯度检查点(gradient checkpointing)(也叫激活重计算)是一种内存优化,而不是容错机制。在前向传播时,你不再为反向传播保存所有激活,而是只保存某些检查点处的激活。反向传播时,再从这些检查点重算缺失的激活。这是用算力换内存:它让前向传播的代价增加约 33%,但可以把激活内存减少 \sqrt{L} 倍(L 是层数)。

  • 把这一切串起来,训练一个前沿模型要把所有这些技术都用上:BF16 混合精度、跨数千 GPU 用环形 all-reduce 的数据并行、节点内张量并行、跨节点流水线并行、用梯度检查点减内存、用 MoE 提升参数效率,外加定期检查点保证容错。这里的系统工程和算法设计一样具有挑战性。

  • 总结一下分布式训练工具箱:

技术 作用 代价
混合精度(BF16) 把激活/梯度内存减半 轻微的数值差异
数据并行 跨 GPU 扩大批量大小 梯度同步的通信开销
张量并行 把层切分到多个 GPU 需要高速互连
流水线并行 把模型阶段分到多个 GPU 流水线气泡(浪费算力)
梯度累积 模拟更大的批量 更慢(多次前向/反向传播)
梯度检查点 减少激活内存 多约 33% 算力
环形 all-reduce 高效的梯度求平均 大模型下受带宽限制
MoE 容量更大,FLOPs 不变 负载均衡、路由复杂
缩放定律 指导算力分配 经验性的,未必在所有尺度上都成立

编程练习(使用 CoLab 或 notebook)

  1. 计算一个 Transformer 层的 FLOPs 和内存需求。给定隐藏维度 d、序列长度 n、批量大小 B 和层数,估算总训练代价。
import jax.numpy as jnp def transformer_layer_flops(d, n, B): """一个 Transformer 层前向传播的近似 FLOPs。""" # QKV 投影:3 * (B * n * d * d) * 2(乘加) qkv_flops = 3 * 2 * B * n * d * d # 注意力:(B * n * n * d) * 2 算 QK^T,(B * n * n * d) * 2 算 attn*V attn_flops = 2 * 2 * B * n * n * d # 输出投影:(B * n * d * d) * 2 out_flops = 2 * B * n * d * d # FFN:两层,d->4d 和 4d->d:2 * (B * n * d * 4d) * 2 ffn_flops = 2 * 2 * B * n * d * 4 * d return qkv_flops + attn_flops + out_flops + ffn_flops def transformer_layer_memory(d, n, B, dtype_bytes=2): """一层激活内存(字节)的近似。""" # QKV:3 * B * n * d qkv_mem = 3 * B * n * d * dtype_bytes # 注意力权重:B * heads * n * n(近似为 B * n * n * sizeof) attn_mem = B * n * n * dtype_bytes # FFN 中间量:B * n * 4d ffn_mem = B * n * 4 * d * dtype_bytes return qkv_mem + attn_mem + ffn_mem # 示例:GPT-2 规模 d, n, B, L = 1024, 1024, 8, 24 fwd_flops = transformer_layer_flops(d, n, B) total_flops = 3 * L * fwd_flops # 3 倍对应前向 + 反向 act_mem = L * transformer_layer_memory(d, n, B) param_count = L * (12 * d * d + 13 * d) # 近似 print(f"Model: d={d}, n={n}, B={B}, L={L}") print(f"Parameters: {param_count / 1e6:.0f}M") print(f"FLOPs per step: {total_flops / 1e12:.2f} TFLOPs") print(f"Activation memory: {act_mem / 1e9:.2f} GB (BF16)") print(f"Parameter memory (FP32): {param_count * 4 / 1e9:.2f} GB") print(f"Adam optimizer memory: {param_count * 8 / 1e9:.2f} GB") print(f"Total training memory: {(param_count * 16 + act_mem) / 1e9:.2f} GB")
  1. 模拟数据并行训练。把数据集切分到多个"虚拟 GPU"上,各自独立计算梯度,求平均,并验证结果与单 GPU 训练一致。
import jax import jax.numpy as jnp # 简单的线性模型:y = wx + b key = jax.random.PRNGKey(0) X = jax.random.normal(key, (64, 4)) w_true = jnp.array([1.0, -2.0, 3.0, 0.5]) y = X @ w_true + 0.1 * jax.random.normal(key, (64,)) def loss_fn(w, X, y): return jnp.mean((X @ w - y) ** 2) grad_fn = jax.grad(loss_fn) # 单 GPU:全批量梯度 w = jnp.zeros(4) grad_single = grad_fn(w, X, y) # 数据并行:切分到 4 个"GPU" n_gpus = 4 chunk_size = len(X) // n_gpus grads = [] for i in range(n_gpus): X_chunk = X[i*chunk_size:(i+1)*chunk_size] y_chunk = y[i*chunk_size:(i+1)*chunk_size] grads.append(grad_fn(w, X_chunk, y_chunk)) # All-reduce:梯度求平均 grad_parallel = jnp.mean(jnp.stack(grads), axis=0) print("Single-GPU gradient:", grad_single) print("Data-parallel gradient (avg):", grad_parallel) print(f"Match: {jnp.allclose(grad_single, grad_parallel, atol=1e-5)}") # 两者都训练并比较 w_single, w_parallel = jnp.zeros(4), jnp.zeros(4) lr = 0.1 for step in range(100): w_single = w_single - lr * grad_fn(w_single, X, y) grads = [grad_fn(w_parallel, X[i*chunk_size:(i+1)*chunk_size], y[i*chunk_size:(i+1)*chunk_size]) for i in range(n_gpus)] avg_grad = jnp.mean(jnp.stack(grads), axis=0) w_parallel = w_parallel - lr * avg_grad print(f"\nAfter 100 steps:") print(f"Single-GPU weights: {w_single}") print(f"Data-parallel weights: {w_parallel}") print(f"Max difference: {jnp.max(jnp.abs(w_single - w_parallel)):.2e}")
  1. 实现一个简单的混合专家层。构造一个把词元路由到 top-K 专家的门控网络,并组合它们的输出。
import jax import jax.numpy as jnp def expert_fn(x, W1, b1, W2, b2): """简单的两层 FFN 专家。""" h = jnp.maximum(0, x @ W1 + b1) # ReLU return h @ W2 + b2 def moe_layer(x, gate_W, experts_params, top_k=2): """ MoE 前向传播。 x: (batch, d_model) gate_W: (d_model, n_experts) experts_params: 每个专家的 (W1, b1, W2, b2) """ n_experts = len(experts_params) # 门控:计算路由分数 gate_logits = x @ gate_W # (batch, n_experts) gate_probs = jax.nn.softmax(gate_logits, axis=-1) # Top-K 选择 top_k_indices = jnp.argsort(-gate_probs, axis=-1)[:, :top_k] top_k_probs = jnp.take_along_axis(gate_probs, top_k_indices, axis=-1) # 重新归一化 top_k_probs = top_k_probs / jnp.sum(top_k_probs, axis=-1, keepdims=True) # 计算各专家输出(简化做法:先全跑,后面再掩码) expert_outputs = jnp.stack([ expert_fn(x, *experts_params[i]) for i in range(n_experts) ], axis=1) # (batch, n_experts, d_model) # 收集 top-K 专家输出并加权 batch_idx = jnp.arange(x.shape[0])[:, None] selected_outputs = expert_outputs[batch_idx, top_k_indices] # (batch, top_k, d_model) output = jnp.sum(selected_outputs * top_k_probs[:, :, None], axis=1) return output, gate_probs # 设置 key = jax.random.PRNGKey(42) batch, d_model, d_ff, n_experts = 8, 16, 32, 4 # 初始化专家 experts_params = [] for i in range(n_experts): k1, k2, key = jax.random.split(key, 3)[0], jax.random.split(key, 3)[1], jax.random.split(key, 3)[2] experts_params.append(( jax.random.normal(k1, (d_model, d_ff)) * 0.1, jnp.zeros(d_ff), jax.random.normal(k2, (d_ff, d_model)) * 0.1, jnp.zeros(d_model), )) key, subkey = jax.random.split(key) gate_W = jax.random.normal(subkey, (d_model, n_experts)) * 0.1 x = jax.random.normal(key, (batch, d_model)) output, gate_probs = moe_layer(x, gate_W, experts_params, top_k=2) print(f"Input shape: {x.shape}") print(f"Output shape: {output.shape}") print(f"Gate probabilities (first sample): {gate_probs[0]}") print(f"Expert usage (avg across batch):") for i in range(n_experts): usage = jnp.mean(gate_probs[:, i]) print(f" Expert {i}: {usage:.3f}")

作者与出处
原作者: HenryNdubuaku
来源:HenryNdubuaku
许可证:Apache-2.0
整理: 灏天文库整理
由灏天文库结构化整理,提供目录导航、全文检索与在线阅读,便于系统化学习
发布者: 作者: HenryNdubuaku 转发
评论区 (0)
U