Lecture 22: Parallel Deep Learning: Data Parallelism
Lecture 22: Parallel Deep Learning: Data Parallelism
1. 章节标题与概述
Lecture 22: Parallel Deep Learning: Data Parallelism(数据并行训练:参数服务器、AllReduce 家族与 ZeRO 零冗余优化器)
本讲核心问题:DNN 训练的本质是”用一批样本算出梯度、把梯度加起来、再统一更新权重“,而这句话里的求和号 $\nabla L(w)=\sum_{j=1}^{n}\nabla L_j(w)$ 天然就是可并行的(讲义 slide 7 直接把 SGD 的更新式写成带 Batch Size 的求和形式)。于是问题变成:这台机器上有 N 张加速卡,怎么把这个”求和”摊到 N 张卡上? 一旦摊开,瓶颈立刻从”算力”变成了”通信“:所有卡必须交换 M 个参数的梯度,而交换方式的不同选择(参数服务器 / Naive / Ring / Tree / Butterfly AllReduce)带来的时间差可以高达 几十倍(slide 25–26)。第二个问题随之而来:数据并行要求每张卡都保存一份完整的模型副本,当模型大到单卡显存装不下时(GPT-3 175B 需要约 2800–3500 GB,slide 29–30),数据并行本身就崩了——这就引出本讲后半部分的 ZeRO(Zero Redundancy Optimizer):把”人人一份”的冗余账本切成 N 份,用通信量换取显存。
- 涉及的主要硬件/软件机制:
- 硬件侧:跨设备互连(卡间 NVLink、主机内 PCIe、机间 InfiniBand / Ethernet)与它们各自的带宽 + 延迟特性;GPU 显存层次(SM 内的寄存器/共享内存 → 6 MB L2 → 16 GB HBM,900 GB/s);集合通信硬件路径(ring、tree、butterfly 三种拓扑对应不同的链路占用模式)。
- 软件侧:mini-batch SGD 与 Adam 更新规则、梯度聚合(gradient aggregation) 的四种算法、AllReduce / AllGather / ReduceScatter 集合通信原语、混合精度训练(FP16 参数与梯度 + FP32 主权重与优化器状态)、ZeRO 的三个阶段(切优化器状态 / 切梯度 / 切参数)、以及把通信与反向计算分桶重叠(bucketing / overlap)的工程手段。
在并行计算知识体系中的角色:这是全课程”从单机并行走向多机并行“的枢纽一讲。前面讲过的算术强度、Roofline、分块、SIMD、warp、共享内存、树的归约,在这里被拼装成一个真实的超大规模负载:计算内部用 GEMM 分块榨干单卡算力,卡之间用 AllReduce 做跨越整个集群的树形/环形归约。它同时是”集合通信”这一主题在课程中的第一次系统登场(Ring / Tree / Butterfly 的通信量与延迟公式就是”归约网络的 work-span 分析”),也是下一讲 Lecture 23(model / pipeline parallelism) 的铺垫:数据并行解决”数据太多”,模型/流水并行解决”模型太大”,而真实训练里两者必须组合(讲义 slide 66 提到 17.2B 的 Turing NLG 就是 Stage 1 + Megatron 张量并行一起上)。
- 配套材料:
lectures/24-parallel_deep_learning_data_parallel.pdf(抽取文本extracted/24-parallel_deep_learning_data_parallel.txt,共 80 页)——已公开,可在公开网络直接下载。Fall 2026 日程表(https://www.cs.cmu.edu/~418/schedule.html)把 Oct 23 排为第 22 讲 “Parallel Deep Learning (data parallelism)”,该行的 slides/video 链接目前以 HTML 注释形式给出(注释原文:”slides/video from a previous offering; uncomment when posted for Fall 2026”),注释中的 slides 指向lectures/24-parallel_deep_learning_data_parallel.pdf,video 指向归档 YouTube 链接 https://www.youtube.com/watch?v=AbsVyQqqIcM。该 PDF 位于公开目录 https://www.cs.cmu.edu/~418/lectures/ 之下。三点说明(不是错误):①文件编号 24 与 Fall 2026 讲次 22 不一致(历史学期讲次重排),讲义首页也写着 “CMU 15-418/15-618, Fall 2025”;②讲义首页署名是 Zhihao Jia(Stanford University),并带有 “Automated Approaches to Accelerate Machine Learning” 的标题条——即这一讲是客座/外请讲座的沿用版本;③讲义中有大量纯动画页(如 slide 42–65 的 ZeRO Stage 1 逐步动画、slide 14–20 的 ring allreduce 逐帧动画)在文本抽取后只剩重复的图注文字,本笔记对这些内容依据同页可见文字(如”Step 1 (Aggregation): each worker send one slice (M/N parameters)…”、”Backward propagation to generate FP16 gradients and AllReduce to average”)展开,不冒充原文中不存在的数字。cs149_supp/dnninference.txt(抽取文本,共 75 页)——已公开。Stanford CS149(Fall 2025)Lecture 9: Efficiently Evaluating DNNs:讲 DNN 前向推理的性能优化(pipelining 与 double buffering、Roofline 与算术强度、循环融合 1/3 → 3/5、分块 GEMM、explicit/implicit GEMM、Flash-Attention 的分块 softmax、低精度 16/8/4-bit、以及 V100 的 6 MB L2 + 900 GB/s HBM)。它是本讲的性能分析工具箱:本讲第 4 节用它给出的算术强度、Roofline、分块思路定量解释”为什么梯度计算可以是计算受限、而 AllReduce 永远只能是带宽受限”。- 讲义中引用的外部材料:Kingma and Ba, “Adam: A Method for Stochastic Optimization”, 2014(https://arxiv.org/abs/1412.6980,slide 33 的 Adam 更新规则出处);Vaswani et al., “Attention is all you need”(slide 34 的 Transformer 结构出处);以及 slide 35–65 标注的 “Adapted from Minjia Zhang, DeepSpeed Presentation”(ZeRO / 内存账本动画的来源)。硬件规格(V100 16G/32G、A100 40G/80G)出自 slide 30。
- 未发布 / 需登录:本讲的讲课录像(Panopto / YouTube)在 Fall 2026 日程表中被注释隐藏,属未发布;Ed 讨论区、Autolab、Canvas 均需登录,非公开。部分讲座(Performance Analysis/Profiling、Transactional Memory、AI in System Design 等)在 Fall 2026 尚未发布讲义,其历史学期 PDF 位于
/afs/cs/academic/class/15418-*/public/之下,需要 CMU 登录,属未公开。 - 课程语境:Fall 2026 授课教师为 Brian Railing 与 Dimitrios Skarlatos;课程由 Kayvon Fatahalian 创建。Fall 2026 日程表把当天标注为 Assignment 4 due;前一讲为 Lecture 21 “Memory Consistency”(Oct 21),后一讲为 Lecture 23 “Parallel Deep Learning (model and pipeline parallelism)”(Oct 26,对应公开 PDF
25-parallel_deep_learning_model_pipeline_parallel.pdf)。
2. 核心概念与硬件/软件架构图解
2.1 回顾:DNN 训练的三个阶段与”可并行性”的来源
定义与目的:讲义 slide 4–6 用三张图把”训练”这件事定义清楚:① 前向传播(forward propagation)——把一批输入样本喂进模型,逐算子计算得到预测;② 反向传播(backward propagation)——把模型”反过来跑”,为每个可训练权重算出一个梯度 $\partial L(w)/\partial w_i$;③ 权重更新(weight update)——用梯度把权重往损失下降的方向推: \(w_i := w_i - \gamma \frac{\partial L(w)}{\partial w_i},\qquad \frac{\partial L(w)}{\partial w_i}=\sum_{j=1}^{n}\frac{\partial l_j(w)}{\partial w_i}\) 其中 $\gamma$ 是学习率(step size),$n$ 是 batch size,求和内的每一项是单个样本的梯度。讲义 slide 7 直接问:”How can we parallelize DNN training?”——答案就写在那个求和号上。
- 直观解释(”它是什么?”):把训练想象成一个复习小组在刷同一本题库。
- 模型是大家共用的一份”解题笔记”;
- 数据是题库里的题目;
- 梯度是”我做完这批题之后,发现笔记里的哪几行该改、改多少”的修改意见;
- SGD 的求和号就是”把组里所有人的修改意见平均一下再一起改笔记”。 关键洞察在于:意见可以各算各的(完全并行、互不干扰),但”平均之后再改”这一步必须所有人改得一模一样——否则下一轮大家手里的笔记就不是同一份了,后面所有计算全错。这就是数据并行的全部张力的来源:计算可以自由拆分,状态必须强制同步。
- 图 1:数据并行的软件执行模型(一个 iteration 的全流程)
一个 iteration:全局 batch = N 张卡 × 每卡 b 个样本
训练集 ┌──────────────────────────────────────────────────────────────────┐
┌─────────┐ │ GPU0 GPU1 GPU2 GPU(N-1) │
│ shard 0 │── b0 ───▶ │ ┌──────────┐ ┌──────────┐ ┌──────────┐ ┌──────────┐ │
│ shard 1 │── b1 ───▶ │ │ 模型副本 │ │ 模型副本 │ │ 模型副本 │ │ 模型副本 │ │
│ shard 2 │── b2 ───▶ │ │ W(全量)│ │ W(全量)│ │ W(全量)│ │ W(全量)│ │
│ ... │ │ └────┬─────┘ └────┬─────┘ └────┬─────┘ └────┬─────┘ │
└─────────┘ │ │ ①forward │ │ │ │
│ │ ②backward │ │ │ │
│ ▼ ▼ ▼ ▼ │
│ g0=∇L(b0) g1=∇L(b1) g2=∇L(b2) g_{N-1} │
│ └───────────────┴────────┬────────┴───────────────┘ │
│ │ ③ AllReduce(SUM) 再除以 N │
│ ▼ ←—— 本讲的主角:通信代价全在这 │
│ ḡ = (1/N)·Σ gᵢ (每张卡都拿到同一个平均梯度) │
│ │ │
│ ▼ ④ W ← W − γ·adam(ḡ) │
│ 每张卡独立算、结果逐位相同 → 副本自然一致 │
└──────────────────────────────────────────────────────────────────┘
★ 可并行部分:① ②(总 Work 固定,机器容量随 N 线性增长 → 理想情况为 N 倍加速)
★ 不可省部分:③ (通信量 ≈ 2M 个参数,与 N 无关地"钉"在每次迭代里 → 加速比天花板)
- 性能特征:前向+反向的计算量随卡数 N 线性分摊(每卡 1/N),而梯度聚合的通信量几乎与 N 无关(下一节会证明 ring allreduce 每卡收发约 2M 个参数)。这意味着数据并行是一个”固定开销 + 可变计算“的结构:卡越多,通信占比越大,最终加速比被”计算时间/通信时间”这个比值封顶。这就是本讲所有性能分析的主线。
2.2 参数服务器(Parameter Server):最直观的集中式方案与它的天花板
定义与目的:讲义 slide 9 给出第一种方案:worker 把梯度 push 给参数服务器(parameter server),再从服务器 pull 回更新后的参数。它解决的问题是”谁来负责把 N 份梯度加起来“——答案是”设一个专门的服务器进程/机器来当那个求和点”,worker 之间不需要互相认识。
直观解释(”它是什么?”):像一个宿舍楼里只有一台公用洗衣机。每个人各自攒一筐衣服(本地梯度),都送到同一间洗衣房,洗好之后再由这台机器把衣服分回给每个人(更新后的参数)。人少的时候很方便、逻辑也简单;但人一多,队伍全堵在那一台机器门口——机器的进水/出水口(服务器网卡带宽)成了唯一通道,排队时间随人数线性上涨。讲义 slide 10 的原话正是:”Centralized communication: all workers communicate with parameter servers for weights update; cannot scale to large numbers of workers“。
图 2:参数服务器 vs AllReduce 的通信拓扑(中心化 vs 去中心化)
(a) Parameter Server:星形,中心带宽 O(N) 增长 (b) Ring AllReduce:环形,只用邻居链路
┌────────────┐ W0 ──▶ W1
┌────▶│ PS / 分片 │◀────┐ ▲ │
│ └────────────┘ │ │ ▼
│ ▲ ▲ │ W3 ◀── W2
┌──┴─┐ ┌──┴─┐ ┌──┴─┐ ┌──┴─┐
│ W0 │ │ W1 │ │ W2 │ │ W3 │ 每个 PS 端口的进出流量 ∝ N 每张卡只与左右邻居说话
└────┘ └────┘ └────┘ └────┘ → 服务器网卡先饱和 每卡收发 2M(N-1)/N ≈ 2M 个参数
push 梯度 + pull 参数(各 M 个) 延迟 ≈ M·N / BW(讲义 slide 26) → 与 N 无关,可扩展
- 性能特征:设模型有 M 个参数、互连带宽 BW,PS 必须串行地服务 N 个 worker 的收发,讲义 slide 26 给出的延迟是 \(T_{PS} \approx \frac{M\cdot N}{\text{bandwidth}}\) 即随 worker 数线性劣化;而且所有 worker 的流量都要挤过服务器这一跳,服务器网卡是唯一瓶颈。讲义紧接着提出问题:”How can we decentralize communication in DNN training?“(slide 10)——答案就是 AllReduce。
2.3 AllReduce 家族:Naive / Ring / Tree / Butterfly
定义与目的:AllReduce 是”逐元素归约(element-wise reduction)”的集合通信原语(slide 11):N 个设备各持一个长度为 M 的向量,操作结束后每个设备都拥有这 N 个向量逐元素之和。它一次解决了两件事:求和(reduce)与把结果广播回来(broadcast)——正好对应数据并行里”平均梯度 + 所有副本保持一致”的需求。
- 直观解释(”它是什么?”):把 AllReduce 想成”全班交换作业本,最后每个人手里都拿到全班的分数合计表“。四种算法的区别只是怎么传本子:
- Naive:每个人都把自己的本子复印 N−1 份、发给其他所有人 → 纸(带宽)消耗 O(N²)。
- Ring(环):大家围成一圈,每人只把一页(M/N 个参数)传给右手边的人,传 N−1 轮,每页沿环被”接力累加”一圈(Aggregation);然后再传一圈把所有页补齐(Broadcast)。像击鼓传花 + 传阅笔记本:每人经手的量固定,队伍再长也不吃亏。
- Tree(树):像公司的组织架构,层层上报(聚合)再层层下发(广播),深度只有 log N,但每个人每次要传整本(M 个参数),而且根节点和上层链路要承受巨大流量。
- Butterfly(蝶形):像每一轮全体换座位两两对接,第 k 轮每个人与”编号差 2^k”的伙伴交换全部 M 个参数并本地求和;log N 轮之后人人拿到总和。轮数最少(低延迟),但每人经手 M·log N 个参数(带宽最贵)。
- 图 3:Ring AllReduce 的 6 步完整时序(N=4,M 个参数切成 c0..c3 四片)
环: W0 ──▶ W1 ──▶ W2 ──▶ W3 ──▶ W0 (每卡只与 left/right 邻居通信)
每片大小 = M/N 个参数;括号内数字 = 该片上已经累加的"贡献数"
阶段 1 Reduce-Scatter(N−1 = 3 步):第 k 步 送 c_{(i−k) mod N},收 c_{(i−k−1) mod N} 并就地把收到的加进自己的副本
步 k=0 送自己的片 W0: c3(2) W1: c0(2) W2: c1(2) W3: c2(2)
步 k=1 送上一轮收到的片 W0: c2(3) W1: c3(3) W2: c0(3) W3: c1(3)
步 k=2 W0: c1(4)✓ W1: c2(4)✓ W2: c3(4)✓ W3: c0(4)✓
★ 结论:worker i 手里"恰好"完成归约的是片 c_{(i+1) mod N}(每片都走满一圈、被加 N−1 次)
阶段 2 AllGather(N−1 = 3 步):第 k 步 送 c_{(i+1−k) mod N}(送出的必须是"已归约片"),从前驱收 c_{(i−k) mod N}
步 k=0 每卡已有 2/4 片 步 k=1 3/4 片 步 k=2 4/4 片 ✓ 人人拥有完整的 c0..c3 之和
总量:2(N−1) = 6 步;每步每卡收发 M/N 个参数
每卡总收发 = 2(N−1)·M/N ≈ 2M 个参数(与 N 无关!)
全系统通信量 = N·2(N−1)·M/N ≈ 2NM(讲义 slide 20 记作 2·M·N)
- 图 4:Tree AllReduce 的组织结构与代价(N=7)
聚合(reduce:向上 log N 步) 广播(broadcast:向下 log N 步)
W0 (根) W0(已持有全体之和)
▲ ▲ ┌─────┴─────┐
┌───────┘ └───────┐ ┌──────┘ └──────┐
W1 W2 W1 W2
▲ ▲ ▲ ▲ ┌──┴──┐ ┌──┴──┐
W3 W4 W5 W6 W3 W4 W5 W6
每步传 M 个参数(整份,不是 M/N) 每步传 M 个参数
步数 = 2 × ⌈log2 7⌉ = 2 × 3 = 6 步(深度小 → 延迟低) 单卡收发量 = 2M·log N ≫ 2M(带宽贵)
全系统通信量 = 2·N·M(讲义 slide 22) 负载不均:根与上层链路是热点
- 图 5:Butterfly AllReduce 的交换模式(N=8,log N = 3 步)
规则:第 k 步(k = 0,1,2),worker i 与 worker (i XOR 2^k) 交换"全部 M 个参数",然后本地相加
步 0(伙伴 = i^1) 步 1(伙伴 = i^2) 步 2(伙伴 = i^4)
0 ◀──▶ 1 2 ◀──▶ 3 0 ◀──▶ 2 1 ◀──▶ 3 0 ◀──▶ 4 1 ◀──▶ 5
4 ◀──▶ 5 6 ◀──▶ 7 4 ◀──▶ 6 5 ◀──▶ 7 2 ◀──▶ 6 3 ◀──▶ 7
└─ 4 对并发 ─┘ └─ 4 对并发 ─┘ └─ 4 对并发 ─┘
→ 3 步后 8 个节点全部持有完整的和
→ 每步每节点收发 M 个参数 ⇒ 单卡 = M·log N,全系统 = N·M·log N(讲义 slide 24)
★ 用"更多带宽"换"更少轮数":延迟 log N 步,但总通信量比 ring 多 log N / 2 倍
(讲义 slide 23 给出的硬件背景:butterfly / omega 多级互连网络,每级做 2×2 交换)
- 性能特征小结(表 1):通信量、单卡通信量、延迟、可扩展性四个维度一起看,才能理解为什么讲义 slide 25 要问”Ring AllReduce 比 Tree / PS 更高效可扩展,为什么?”
表 1:四种梯度聚合方案的定量对比(M = 参数量,N = worker 数,BW = 链路带宽,α = 单次消息启动延迟)
| 方案 | 全系统通信量 | 单卡收发量 | 延迟(讲义 slide 26 的带宽项) | 完整 α-β 延迟模型 | 负载均衡 | 可扩展性 |
|---|---|---|---|---|---|---|
| Parameter Server | 2·N·M | 2M(但服务器侧 ∝ N·M) | M·N / BW | N·M/BW + 2α | 差(服务器是热点) | 差:中心带宽随 N 线性劣化 |
| Naïve AllReduce | N²·M(= N(N−1)M) | (N−1)M | (N−1)M / BW | (N−1)·(α + M·w) | 好 | 差:O(N²) 通信 |
| Ring AllReduce | 2·N·M | ≈2M(与 N 无关) | 2M / BW | 2(N−1)·(α + (M/N)·w) | 好 | 好:单卡量恒定 |
| Tree AllReduce | 2·N·M | 2M·log N | 2·log N·M / BW | 2·log N·(α + M·w) | 差(根/上层热点) | 中:延迟低、带宽贵 |
| Butterfly AllReduce | N·M·log N | M·log N | log N·M / BW | log N·(α + M·w) | 好 | 中:轮数最少、带宽最贵 |
注:表中”全系统通信量”一列取自讲义 slide 25 的原始表述(PS/Naive/Ring/Tree/Butterfly 分别为 $2NM$、$N^2M$、$2NM$、$2NM$、$NM\log N$);”延迟”一列取自 slide 26(PS: $MN/BW$;Ring: $(M/N)\cdot 2N/BW$;Tree: $M\cdot 2\log N/BW$)。“完整 α-β 延迟模型”那一列是本笔记补充的(讲义的延迟公式只写了带宽项,没有显式的 α 项),其中 $w$ = 每个参数元素的传输时间(= 每参数字节数 / BW,fp16 且链路 10 GB/s 时 $w = 5\times10^{-11}$ s),$\alpha$ = 单次消息的启动延迟。这一列把”步数多但每步小”和”步数少但每步大”的差别显性化,4.4 节用它算出 cross-over 点。
2.4 数据并行的内存墙:混合精度训练的”内存账本”
- 定义与目的:数据并行的第一个致命假设是”每张卡都存一份完整模型“(讲义 slide 28:Each GPU saves a replica of the entire model),因此模型参数超过单卡显存就完全无法训练。讲义 slide 29–30 用一张表把这个问题量化,并用一句 “Out of Memory” 收尾(V100 只有 16G/32G,A100 是 40G/80G)。
表 2:大模型规模增长(讲义 slide 29 原始数据)
| 模型 | Parameters | Layers | Hidden Dim | Relative Computation | Memory Footprint |
|---|---|---|---|---|---|
| Bert-Large | 0.32B | 24 | 1024 | 1× | 5.12 GB |
| GPT-2 | 1.5B | 48 | 1600 | 4.7× | 24 GB |
| Turing NLG 17.2B | 17.2B | 78 | 4256 | 54× | 275 GB |
| GPT-3 | 175B | 96 | 12288 | 547× | 2800 GB |
一致性校验(本笔记的推算,用来确认读懂了口径):表中的 Memory Footprint 恰好等于 16 字节/参数(0.32B×16 = 5.12 GB;17.2B×16 = 275.2 GB;175B×16 = 2800 GB)。这 16 字节 = FP32 主权重(4) + FP32 梯度(4) + Adam 一阶动量(4) + Adam 二阶动量(4)。而 slide 40 给出的混合精度训练实际账本是 20 字节/参数——多出来的 4 字节是”额外保留一份 FP16 参数与 FP16 梯度”(2+2)。两个口径不要混用。
直观解释(”它是什么?”):混合精度训练就像一套”正式账本 + 便签”的记账方式。计算(前向/反向)用 FP16,因为算得快、占地方小;但直接拿 FP16 累加更新会因舍入误差把模型”记花”,所以必须再留一份 FP32 的正式账本(master weights),外加 Adam 需要的两个 FP32 统计量(动量、方差)。于是每 1 个参数,要在显存里占 20 个字节——1B 参数的模型就是 20 GB/卡(讲义 slide 40 的原话:”Example 1B parameter model -> 20GB/GPU”,并特别注明”Memory consumption doesn’t include: Input batch + activations”)。
图 6:单卡显存账本与 ZeRO 三级切割
每张卡上的"训练状态"账本(M = 参数量;单位:字节/参数) 1B 参数模型 → 20 GB/卡
┌─────────────────────────────────────────────────────────┐
│ FP16 参数(前向/反向实际使用的副本) 2 B │ ← 便签
│ FP16 梯度 2 B │
├─────────────────────────────────────────────────────────┤
│ FP32 master 参数 4 B │ ┐
│ FP32 梯度 4 B │ ├ 16 B = "FP32 优化器状态"
│ FP32 Adam 一阶动量 m 4 B │ │ (讲义 slide 40:16M bytes)
│ FP32 Adam 二阶动量 v 4 B │ ┘
├─────────────────────────────────────────────────────────┤
│ 合计 20 B │
│ ★ 未计入:输入 batch 与全部中间激活 activations │
└─────────────────────────────────────────────────────────┘
ZeRO 把这份账本按 N 张卡"切"开(下表为每卡显存,单位 M bytes,N = 卡数)
基线数据并行 : [2 参数][2 梯度][ 16 优化器状态 ] = 20M ← 每卡一份全量,卡数不省内存
Stage 1 : [2 参数][2 梯度][ 16/N 优化器状态 ] = 4M + 16M/N
Stage 2 : [2 参数][ (2+16)/N ] = 2M + 18M/N
Stage 3 : [ (2+2+16)/N ] = 20M/N ← 参数也切了
★ 切得越狠,省得越多,但每级都要额外付出一次集合通信(ReduceScatter / AllGather)
- 性能特征:单纯看计算,FP16 让 Tensor Core 吞吐翻数倍;但真正的系统收益在于显存:M 个参数的账本从 20M bytes 变成
20M/Nbytes(Stage 3),代价是每张卡每次前向要 AllGather 别人的参数、反向完要 ReduceScatter 梯度(讲义 slide 76–78:GPUs broadcast their parameters during forward / Parameters are discarded right after use / GPUs broadcast their parameters again during backward)。
2.5 ZeRO:Zero Redundancy Optimizer(把”人人一份”变成”人人一份之 1/N”)
定义与目的:讲义 slide 31 定义 ZeRO 为”Eliminating data redundancy in data parallel training“,并说它是”a widely used technique for data parallel training of large models“。它要解决的就是 2.4 节的内存墙:既然 N 张卡上存了 N 份完全相同的优化器状态/梯度/参数,那就别存了——每张卡只负责其中 1/N,需要时再向别人要。
直观解释(”它是什么?”):像一套百科全书的联机共享。过去是每个宿舍都买一整套(贵、且 90% 的时间在落灰);ZeRO 改成各分馆分藏不同卷,谁要查哪一卷就快递过来,用完立刻寄回去(讲义 slide 77 的 “Parameters are discarded right after use”)。查得频繁,快递费(通信)就涨——这就是 ZeRO 用通信量换显存的交换关系。 三个阶段可以类比”先共享最贵的、再共享次贵的”:Stage 1 先共享最占地方、用得最少的(优化器状态,16 B/参数的巨无霸);Stage 2 再共享梯度;Stage 3 连参数也共享。
图 7:ZeRO 三个阶段的”切什么”与流水线动作
┌─────────────────────── Stage 1:切优化器状态(Partitioning Optimizer States)────────────────────┐
│ 基线:每卡 [参数 2][梯度 2][优化器状态 16] = 20M │
│ 现在:每卡 [参数 2][梯度 2][优化器状态 16/N] ← 每卡只更新自己那一份参数! │
│ 一个 iteration 的动作序列(讲义 slide 42–65 的动画文字): │
│ ① forward 跑完整个 transformer 栈 (每卡都是完整模型,算自己的数据) │
│ ② backward 得到 FP16 梯度 │
│ ③ AllReduce 求平均梯度 │
│ ④ 每卡用 Adam 更新"自己负责的那 1/N 个参数的" FP32 master 权重 │
│ ⑤ 同步出对应的 FP16 权重 │
│ ⑥ AllGather FP16 权重,补齐整份模型 → 回到 ① │
└──────────────────────────────────────────────────────────────────────────────────────────────┘
┌─────────────────────── Stage 2:再切梯度(Partitioning Gradients)─────────────────────────────┐
│ 每卡 [参数 2][梯度+优化器状态 (2+16)/N] │
│ 关键工程动作(讲义 slide 68–72):*Perform AllReduce right after back propagation of each layer* │
│ → 不是等整个反向跑完再 AllReduce,而是"每层反完就立刻 Reduce",只有"负责更新该参数的卡"保留梯度 │
│ → 好处:① 梯度缓冲只有 1/N ② 通信可以与该层之后的反向计算重叠(这就是 bucketing/overlap) │
└──────────────────────────────────────────────────────────────────────────────────────────────┘
┌─────────────────────── Stage 3:连参数也切(Partitioning Parameters)──────────────────────────┐
│ 每卡 [ (2+2+16)/N ] = 20M/N ← 单卡显存真正随 N 线性下降 │
│ 代价(讲义 slide 74–78): │
│ forward 时 GPUs broadcast their parameters(用到哪层就 AllGather 哪层) │
│ 用完立刻丢弃(discard right after use)→ 省显存、但参数会被反复搬运 │
│ backward 时 GPUs broadcast their parameters again │
│ → 通信量从基线 DP 的 2M 涨到 3M(多出"前向 + 反向各一次参数 AllGather") │
└──────────────────────────────────────────────────────────────────────────────────────────────┘
- 性能特征:讲义 slide 66 / 79 强调 ZeRO 是”progressive memory savings and communication volume“(逐步省内存、通信量逐步上升),并给出一个真实案例:17.2B 的 Turing NLG 由 Stage 1 + Megatron 支撑(即 ZeRO-1 + 张量并行组合使用,而不是只用一种并行方式)。
表 3:ZeRO 三阶段的内存与通信账本(M = 参数量,N = 卡数;通信量为”每卡每个 iteration 搬运的参数元素数”)
| 配置 | 每卡显存(bytes) | 175B 模型、N=64 时的每卡显存 | 每卡通信量(元素) | 通信量(fp16 字节) | 相对基线通信 |
|---|---|---|---|---|---|
| 基线数据并行 | 20M | 3500 GB(不可能) | 2M(AllReduce 梯度) | 4M = 700 GB | 1.0× |
| ZeRO Stage 1 | 4M + 16M/N | 743.8 GB(仍不可能) | 2M(ReduceScatter 梯度 + AllGather 参数) | 700 GB | 1.0× |
| ZeRO Stage 2 | 2M + 18M/N | 399.2 GB(仍不可能) | 2M(同上,但梯度被切、可重叠) | 700 GB | 1.0× |
| ZeRO Stage 3 | 20M/N | 54.7 GB(A100-80GB 装得下) | 3M(两次参数 AllGather + 一次梯度 ReduceScatter) | 1050 GB | 1.5× |
读表要点:Stage 1/2 的通信量与基线完全相同(这是 ZeRO 论文里最反直觉的结论之一:把 AllReduce 拆成 ReduceScatter + AllGather,总量不变),只有 Stage 3 因为要多传一轮参数而变成 1.5 倍。也就是说 “省显存”在前两个阶段几乎是免费的,第三个阶段才要付出 50% 的通信代价。表中 175B 的显存数字由
20M/N、4M+16M/N、2M+18M/N代入 M = 175×10⁹ 直接算出(1 GB = 10⁹ bytes)。
2.6 硬件视角:这些通信到底跑在什么上
定义与目的:AllReduce 的代价最终由物理链路决定。单机内 GPU 之间走 NVLink(V100 上 6 条链路 × 25 GB/s/方向),跨主机走 PCIe + 网卡(InfiniBand / Ethernet);而单卡内部的算力则由 GPU 的内存层次决定(CS149 slide 48 的 V100 剖面:80 个 SM、6 MB L2、16 GB HBM、900 GB/s)。
直观解释(”它是什么?”):把集群想成一个城市的物流网:SM 内部是”货架到工作台”的距离(寄存器/共享内存),900 GB/s 的 HBM 是”市内主干道”,NVLink 是”楼内电梯”,InfiniBand 是”城际高速”。AllReduce 是必须跑满全城的重卡运输——所以数据并行训练的性能几乎总由”城际高速的口径”决定,而不是由”工作台算得多快”决定。
图 8:GPU 内存层次与互连(硬件结构;数值取自 CS149 slide 48 与讲义 slide 30)
┌──────────────────────────── SM(共 80 个)────────────────────────────┐
│ 寄存器文件(每 SM 256 KB)← 分块 GEMM 的 C 累加器常驻,带宽最高 │
│ ├── 64 个 FP32 lane / warp scheduler × 4 │
│ ├── Tensor Core(FP16/BF16 4×4×4 矩阵乘加)→ V100 FP16 ≈ 125 TFLOP/s │
│ └── 共享内存 / L1(可配 ~96 KB)← implicit GEMM 的 tile 暂存处 │
├──────────────────────────── ... 其余 79 个 SM ... ─────────────────────┤
│ L2 Cache 6 MB(全芯片共享) │
├───────────────────────────────────────────────────────────────────────┤
│ HBM 显存 16 GB(V100)/ 40~80 GB(A100) 带宽 900 GB/s │
└───────────────────────────────────────────────────────────────────────┘
▲ ▲ ▲
│ 卡内:算力 ↔ 900 GB/s 的博弈 │ 卡间:NVLink 2.0 │ 机间:IB / Ethernet
│ (决定单卡能不能"喂饱") │ ≈ 150 GB/s/方向 │ 100 Gb/s ≈ 12.5 GB/s
│ │ PCIe Gen3 x16 ≈ 12.6 │ 400 Gb/s ≈ 50 GB/s
★ 关键推论:梯度缓冲是"每卡一份全量 M 个参数",它【必须】穿过最外面那层最细的管子
→ 单卡算力越强、链路越细,数据并行的扩展性就越差
- 性能特征:这张图直接给出两条性能约束:
- 算力约束:单卡要跑满 125 TFLOP/s(FP16 Tensor Core),数据复用必须做到 125e12/900e9 ≈ 139 FLOP/byte 以上(V100 的 ridge point,CS149 slide 6–8 的 Roofline 拐点);
- 通信约束:每一次迭代都必须搬运 2M 个参数穿越 NVLink 或网卡,这个量不随卡数减少。
3. 代码示例与性能分析
3.1 示例一:手写 MPI Ring AllReduce(把讲义的 6 步算法变成可跑的程序)
- 代码
/* ring_allreduce.c —— 手写 Ring AllReduce,并与 MPI_Allreduce 对拍校验
* 编译: mpicc -O3 -march=native -std=c11 ring_allreduce.c -o ring_allreduce -lm
* 运行: mpirun -np 8 --oversubscribe ./ring_allreduce 100000000 5
* 参数 1 = 每张卡本地梯度的元素数 M;参数 2 = 重复次数(取最快一次)
*/
#include <mpi.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <math.h>
#include <limits.h>
/* 就地做一次 Ring AllReduce:recvbuf 返回"所有 rank 的 sendbuf 之和"
* 通信模式与讲义 slide 14-20 完全一致:ReduceScatter(N-1 步) + AllGather(N-1 步) */
static double ring_allreduce(const float *sendbuf, float *recvbuf, long M,
int rank, int size, float *scratch)
{
const long chunk = M / size; /* 每片 M/N 个参数 */
const int left = (rank - 1 + size) % size; /* 环上前驱 */
const int right = (rank + 1) % size; /* 环上后继 */
memcpy(recvbuf, sendbuf, (size_t)M * sizeof(float));
const double t0 = MPI_Wtime();
/* ---------- 阶段 1:Reduce-Scatter(N-1 步,就地累加)---------- */
for (int k = 0; k < size - 1; k++) {
const int send_c = (rank - k + size) % size; /* 本步交出哪一片 */
const int recv_c = (rank - k - 1 + size) % size; /* 本步收到哪一片 */
MPI_Sendrecv(recvbuf + (size_t)send_c * chunk, (int)chunk, MPI_FLOAT, right, 0,
scratch, (int)chunk, MPI_FLOAT, left, 0,
MPI_COMM_WORLD, MPI_STATUS_IGNORE);
float *dst = recvbuf + (size_t)recv_c * chunk;
for (long i = 0; i < chunk; i++) dst[i] += scratch[i]; /* 本地累加 = reduce */
}
/* ---------- 阶段 2:AllGather(N-1 步,把已归约片传一圈)---------- */
for (int k = 0; k < size - 1; k++) {
const int send_c = (rank + 1 - k + size) % size; /* 送"刚拿到的已归约片" */
const int recv_c = (rank - k + size) % size;
MPI_Sendrecv(recvbuf + (size_t)send_c * chunk, (int)chunk, MPI_FLOAT, right, 1,
recvbuf + (size_t)recv_c * chunk, (int)chunk, MPI_FLOAT, left, 1,
MPI_COMM_WORLD, MPI_STATUS_IGNORE);
}
return MPI_Wtime() - t0;
}
int main(int argc, char **argv)
{
MPI_Init(&argc, &argv);
int rank, size;
MPI_Comm_rank(MPI_COMM_WORLD, &rank);
MPI_Comm_size(MPI_COMM_WORLD, &size);
const long M = (argc > 1) ? atol(argv[1]) : 1000000L;
const int R = (argc > 2) ? atoi(argv[2]) : 3;
if (M % size != 0 || M / size > INT_MAX) {
if (rank == 0) fprintf(stderr, "要求 M 能被 size 整除,且 M/size <= INT_MAX\n");
MPI_Finalize();
return 1;
}
float *local = (float *)malloc((size_t)M * sizeof(float));
float *mine = (float *)malloc((size_t)M * sizeof(float));
float *ref = (float *)malloc((size_t)M * sizeof(float));
float *scratch = (float *)malloc((size_t)(M / size) * sizeof(float));
for (long i = 0; i < M; i++) /* 每卡不同的"本地梯度" */
local[i] = (float)(((rank * 131 + i) % 17) - 8) * 0.25f;
double best = 1e30;
for (int r = 0; r < R; r++) { /* 计时:取最快一次,排除冷启动 */
MPI_Barrier(MPI_COMM_WORLD);
const double t = ring_allreduce(local, mine, M, rank, size, scratch);
if (t < best) best = t;
}
MPI_Allreduce(local, ref, (int)M, MPI_FLOAT, MPI_SUM, MPI_COMM_WORLD); /* 参考实现 */
double myerr = 0.0;
for (long i = 0; i < M; i++) {
const double e = fabs((double)mine[i] - (double)ref[i]);
if (e > myerr) myerr = e;
}
double gerr = 0.0;
MPI_Reduce(&myerr, &gerr, 1, MPI_DOUBLE, MPI_MAX, 0, MPI_COMM_WORLD);
/* 每卡"发出 + 收到"的总字节数:2(N-1) 步 × 每步 2 次传输 × chunk 个 float */
const double bytes = 2.0 * (size - 1) * (double)(M / size) * sizeof(float) * 2.0;
if (rank == 0) {
printf("N=%2d M=%ld(每卡 %.1f MB) 单步片大小 = %.1f KB\n",
size, M, (double)M * 4 / 1e6, (double)(M / size) * 4 / 1e3);
printf(" Ring AllReduce 时间 = %.3f ms;每卡收发 %.1f MB;有效带宽 = %.2f GB/s\n",
best * 1e3, bytes / 1e6, bytes / best / 1e9);
printf(" 与 MPI_Allreduce 的最大逐元素误差 = %.3e\n", gerr);
printf(" 理论上每卡收发 2*M*(N-1)/N 个参数 = %.1f MB(与上式一致)\n",
2.0 * M * (size - 1) / size * sizeof(float) / 1e6);
}
free(local); free(mine); free(ref); free(scratch);
MPI_Finalize();
return 0;
}
- 【代码做什么?】(逐步解释执行流程)
- 初始化与切分:进入
ring_allreduce后先memcpy一份本地梯度到recvbuf(这一步对应”每张卡都持有完整模型和本地梯度”),然后把 M 个参数按 rank 切成chunk = M/N大小的 N 片,并算出环上的前驱left、后继right。只要邻居,不要中心——这正是讲义 slide 10 说的 “decentralize communication”。 - 阶段 1(Reduce-Scatter,N−1 步):第 k 步,rank 把第
(rank-k) mod N片发给后继,同时从前驱收第(rank-k-1) mod N片,就地累加进自己的recvbuf。执行完 N−1 步之后,按 2.3 节的推导,rank i 手里恰好有一片是完整的和(片(i+1) mod N)——这就是图 3 里”w0: c1(4)✓”那一行。 - 阶段 2(AllGather,N−1 步):第 k 步发送
(rank+1-k) mod N片(必须是这一步手里那一片”已归约”的片),接收(rank-k) mod N片并直接覆盖。因为每步收到的片在下一步才需要转发,覆盖是安全的——这正是流水线的关键:接收缓冲区同时是下一步的发送缓冲区,不需要双倍缓冲。 - 校验与计时:
MPI_Allreduce作为参考实现,逐元素比最大误差(浮点求和不满足结合律,因此不能要求逐位相等);计时取 R 次最快值,并用bytes = 4(N−1)·chunk反算”有效带宽”,和理论值2M(N−1)/N对拍。
- 初始化与切分:进入
- 【并行机制与性能解说】
- 谁在并行、怎么分工:这里并行的是 N 个 MPI 进程(通常跨机器),每个进程内部的累加循环是纯串行的、SIMD 友好的(
dst[i] += scratch[i]会被编译器向量化成vaddps)。进程间的”工作分配”由环下标算术(rank ± k) mod N静态决定,没有任何动态调度 —— 这也是 ring 能做到”负载完美均衡”的原因。 - 共享数据怎么处理:没有共享内存;每个进程持有 M 个参数的私有副本,跨进程传递靠显式的
MPI_Sendrecv(点对点、双向同时进行,避免死锁),”共享”只发生在阶段 1 的逐元素加法这一瞬(每一步只交换 M/N 片)。 - Work / Span / 并行度(用第 2.3 节的算法结构直接推):
- Work(浮点加法总数) = 每卡 (N−1)·(M/N) 次加法 × N 卡 = M(N−1) 次加法;
- Span(关键路径) = 单个参数的累加链长度 = N−1 次相关加法;用时间计则还要加上通信步数:2(N−1) 轮,每轮耗时 α + (M/N)·w(α = 消息启动延迟,w = 每元素传输时间);
- 算法内禀并行度 = Work / Span = M(N−1)/(N−1) = M —— 也就是说参数维度是完全独立的,这是”数据并行能扩展”的根本原因;
- 但单张卡在纯 ring 里只能用上 M/N 路并行(每片内 M/N 个独立加法 + 每片一条累加链)。这就是为什么 N 很大时单卡算力吃不满、真实库(NCCL)要把梯度切成多个并发环(channels)才能同时打满多块网卡/多条 NVLink。
- 瓶颈分析:
- 带宽瓶颈(主导):每卡至少要把 ≈2M 个参数(收发合计;精确值 $2M(N-1)/N$)穿过链路。例:M = 1.5×10⁹ 参数、fp16 梯度 = 3.0 GB,则每卡收发 = $2(N-1)/N \times 3.0$ GB,N=8 时是 5.25 GB(N 很大时趋近上界 2M 个参数 ≈ 6.0 GB);在 12.6 GB/s 的 PCIe Gen3 x16 上需要 ≈ 417 ms,而在 150 GB/s 的 NVLink 上只要 35 ms——12 倍差距。
- 延迟瓶颈(小模型/大 N 时主导):步数 2(N−1) 随 N 线性增长。M = 10⁶(fp16 只有 2 MB)时,每步的传输时间 (M/N)/BW 可能比 α(几微秒)还小,通信完全由 2(N−1)·α 决定。
- 浮点求和的非确定性:ring 的累加顺序随 N 变化,因此换卡数就会改变数值结果的最低几位——校验时只能用容差而不是
==(这也解释了两个版本文档里gerr只打印数量级)。 - 无伪共享问题:这个程序里没有任何共享缓存行(跨机通信),但阶段 1 的
dst[i] +=循环值得注意——它是流式访问,其性能上限是内存带宽而不是算力(M/N 片 ≥ 缓存时)。
- 谁在并行、怎么分工:这里并行的是 N 个 MPI 进程(通常跨机器),每个进程内部的累加循环是纯串行的、SIMD 友好的(
3.2 示例二:OpenMP 小批量 SGD —— “逐样本外积累加” vs “batch 维 GEMM”
- 代码
// mlp_sgd.cpp —— 两层 MLP 的小批量 SGD:对比两种权重梯度算法
// A) 逐样本反向传播 + 每线程私有梯度缓冲 + critical 归约(直观、但访存爆炸)
// B) 先算完整批激活,再用"batch 维 GEMM"算 dW = G^T·X(分块复用、算术强度高)
// 编译: g++ -O3 -fopenmp -march=native -std=c++17 mlp_sgd.cpp -o mlp_sgd
// 运行: OMP_NUM_THREADS=16 ./mlp_sgd 256 784 2048 10 5
// 参数: B(batch) D0(输入维) D1(隐层) D2(类别数) 迭代次数
#include <omp.h>
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <cmath>
#include <vector>
#include <algorithm>
#include <chrono>
using Clock = std::chrono::steady_clock;
static double sec_since(Clock::time_point t0) {
return std::chrono::duration<double>(Clock::now() - t0).count();
}
/* ================= 版本 A:逐样本反向 + 线程私有梯度 + critical 归约 ================= */
static void grads_per_sample(const std::vector<float>& X, const std::vector<float>& Y,
const std::vector<float>& W1, const std::vector<float>& b1,
const std::vector<float>& W2, const std::vector<float>& b2,
int B, int D0, int D1, int D2,
std::vector<float>& dW1, std::vector<float>& db1,
std::vector<float>& dW2, std::vector<float>& db2)
{
std::fill(dW1.begin(), dW1.end(), 0.f);
std::fill(db1.begin(), db1.end(), 0.f);
std::fill(dW2.begin(), dW2.end(), 0.f);
std::fill(db2.begin(), db2.end(), 0.f);
#pragma omp parallel
{
// 每个线程一份"完整的"梯度副本:这就是数据并行里"每卡一份全量"的同构问题
std::vector<float> lW1((size_t)D1 * D0, 0.f), lW2((size_t)D2 * D1, 0.f);
std::vector<float> lb1(D1, 0.f), lb2(D2, 0.f);
std::vector<float> h(D1), dlog(D2);
#pragma omp for schedule(static)
for (int n = 0; n < B; n++) { // ← 样本维:天然并行、无需同步
const float* x = &X[(size_t)n * D0];
/* ---- 前向:h = relu(W1 x + b1) ---- */
for (int j = 0; j < D1; j++) {
float s = b1[j];
const float* w = &W1[(size_t)j * D0];
for (int i = 0; i < D0; i++) s += w[i] * x[i];
h[j] = s > 0.f ? s : 0.f;
}
/* ---- 输出层 + softmax 交叉熵:dlogits = p - onehot(y) ---- */
float mx = -1e30f, sum = 0.f;
for (int c = 0; c < D2; c++) {
float s = b2[c];
const float* w = &W2[(size_t)c * D1];
for (int j = 0; j < D1; j++) s += w[j] * h[j];
dlog[c] = s;
if (s > mx) mx = s;
}
for (int c = 0; c < D2; c++) { dlog[c] = expf(dlog[c] - mx); sum += dlog[c]; }
for (int c = 0; c < D2; c++) dlog[c] /= sum;
dlog[(int)Y[n]] -= 1.0f;
/* ---- 反向:外积累加 dW2 += dlog ⊗ h,dW1 += g ⊗ x ---- */
for (int c = 0; c < D2; c++) {
const float g = dlog[c];
lb2[c] += g;
float* lw = &lW2[(size_t)c * D1];
for (int j = 0; j < D1; j++) lw[j] += g * h[j];
}
for (int j = 0; j < D1; j++) {
float g = 0.f;
for (int c = 0; c < D2; c++) g += W2[(size_t)c * D1 + j] * dlog[c];
g = (h[j] > 0.f) ? g : 0.f; // relu 的导数
lb1[j] += g;
float* lw = &lW1[(size_t)j * D0];
for (int i = 0; i < D0; i++) lw[i] += g * x[i];
}
}
/* ---- 合并:这就是节点内的 "AllReduce",代价是 critical 串行区 ---- */
#pragma omp critical
{
for (size_t i = 0; i < dW1.size(); i++) dW1[i] += lW1[i];
for (size_t i = 0; i < dW2.size(); i++) dW2[i] += lW2[i];
for (int j = 0; j < D1; j++) db1[j] += lb1[j];
for (int c = 0; c < D2; c++) db2[c] += lb2[c];
}
}
}
/* ================= 版本 B:batch 维 GEMM(先存激活,再分块算 dW)================= */
static void grads_batched(const std::vector<float>& X, const std::vector<float>& Y,
const std::vector<float>& W1, const std::vector<float>& b1,
const std::vector<float>& W2, const std::vector<float>& b2,
int B, int D0, int D1, int D2,
std::vector<float>& dW1, std::vector<float>& db1,
std::vector<float>& dW2, std::vector<float>& db2)
{
std::vector<float> H((size_t)B * D1), G1((size_t)B * D1), G2((size_t)B * D2);
#pragma omp parallel for schedule(static)
for (int n = 0; n < B; n++) { // 前向:batch 维并行
const float* x = &X[(size_t)n * D0];
for (int j = 0; j < D1; j++) {
float s = b1[j];
const float* w = &W1[(size_t)j * D0];
for (int i = 0; i < D0; i++) s += w[i] * x[i];
H[(size_t)n * D1 + j] = s > 0.f ? s : 0.f; // 激活必须留到反向用(内存换复用)
}
float mx = -1e30f, sum = 0.f;
for (int c = 0; c < D2; c++) {
float s = b2[c];
const float* w = &W2[(size_t)c * D1];
for (int j = 0; j < D1; j++) s += w[j] * H[(size_t)n * D1 + j];
G2[(size_t)n * D2 + c] = s;
if (s > mx) mx = s;
}
for (int c = 0; c < D2; c++) { G2[(size_t)n*D2+c] = expf(G2[(size_t)n*D2+c] - mx); sum += G2[(size_t)n*D2+c]; }
for (int c = 0; c < D2; c++) G2[(size_t)n * D2 + c] /= sum;
G2[(size_t)n * D2 + (int)Y[n]] -= 1.0f;
for (int j = 0; j < D1; j++) {
float g = 0.f;
for (int c = 0; c < D2; c++) g += W2[(size_t)c * D1 + j] * G2[(size_t)n * D2 + c];
G1[(size_t)n * D1 + j] = (H[(size_t)n * D1 + j] > 0.f) ? g : 0.f;
}
}
std::fill(dW1.begin(), dW1.end(), 0.f);
std::fill(db1.begin(), db1.end(), 0.f);
std::fill(dW2.begin(), dW2.end(), 0.f);
std::fill(db2.begin(), db2.end(), 0.f);
/* dW1 = G1^T · X:并行在"权重行"上(每线程独占若干行 → 无竞争、无伪共享),
行内以 batch 为累加轴 → X 的每一行被 D1 个输出通道复用(复用率 = D1) */
#pragma omp parallel for schedule(static)
for (int j = 0; j < D1; j++) {
float* row = &dW1[(size_t)j * D0];
float bsum = 0.f;
for (int n = 0; n < B; n++) {
const float g = G1[(size_t)n * D1 + j];
const float* x = &X[(size_t)n * D0];
bsum += g;
for (int i = 0; i < D0; i++) row[i] += g * x[i];
}
db1[j] = bsum;
}
/* dW2 = G2^T · H */
#pragma omp parallel for schedule(static)
for (int c = 0; c < D2; c++) {
float* row = &dW2[(size_t)c * D1];
float bsum = 0.f;
for (int n = 0; n < B; n++) {
const float g = G2[(size_t)n * D2 + c];
const float* h = &H[(size_t)n * D1];
bsum += g;
for (int j = 0; j < D1; j++) row[j] += g * h[j];
}
db2[c] = bsum;
}
}
int main(int argc, char** argv)
{
const int B = (argc > 1) ? atoi(argv[1]) : 256;
const int D0 = (argc > 2) ? atoi(argv[2]) : 784;
const int D1 = (argc > 3) ? atoi(argv[3]) : 2048;
const int D2 = (argc > 4) ? atoi(argv[4]) : 10;
const int IT = (argc > 5) ? atoi(argv[5]) : 5;
std::vector<float> X((size_t)B*D0), Y(B), W1((size_t)D1*D0), b1(D1, 0.f),
W2((size_t)D2*D1), b2(D2, 0.f);
srand(1234);
auto rnd = []{ return (float)rand() / RAND_MAX - 0.5f; };
for (auto& v : X) v = rnd();
for (auto& v : W1) v = rnd() * 0.05f;
for (auto& v : W2) v = rnd() * 0.05f;
for (int n = 0; n < B; n++) Y[n] = (float)(rand() % D2);
std::vector<float> dW1a((size_t)D1*D0), db1a(D1), dW2a((size_t)D2*D1), db2a(D2);
std::vector<float> dW1b((size_t)D1*D0), db1b(D1), dW2b((size_t)D2*D1), db2b(D2);
grads_per_sample(X, Y, W1, b1, W2, b2, B, D0, D1, D2, dW1a, db1a, dW2a, db2a);
grads_batched (X, Y, W1, b1, W2, b2, B, D0, D1, D2, dW1b, db1b, dW2b, db2b);
double err = 0.0; // 两版本必须算出同一个梯度(容差比较)
for (size_t i = 0; i < dW1a.size(); i++) err = std::max(err, (double)fabsf(dW1a[i] - dW1b[i]));
auto t0 = Clock::now();
for (int it = 0; it < IT; it++) grads_per_sample(X, Y, W1, b1, W2, b2, B, D0, D1, D2, dW1a, db1a, dW2a, db2a);
const double sA = sec_since(t0) / IT;
auto t1 = Clock::now();
for (int it = 0; it < IT; it++) grads_batched(X, Y, W1, b1, W2, b2, B, D0, D1, D2, dW1b, db1b, dW2b, db2b);
const double sB = sec_since(t1) / IT;
const double flops = 6.0 * B * (double)D0 * D1; // 前向 2 + 反向 4(按主导层估)
const double w1b = (double)D1 * D0 * 4; // dW1 的字节数
const double trA = (double)B * (w1b + 2 * w1b); // 版本 A:读 W1 + 读写 dW1
const double trB = ((double)B*D0 + 2.0*B*D1 + (double)D1*D0) * 4; // 版本 B
printf("线程数=%d B=%d D0=%d D1=%d D2=%d\n", omp_get_max_threads(), B, D0, D1, D2);
printf("版本 A(逐样本外积累加): %8.2f ms %7.1f GFLOP/s 估算 DRAM 流量 %6.2f GB AI≈%.2f FLOP/byte\n",
sA * 1e3, flops / sA / 1e9, trA / 1e9, flops / trA);
printf("版本 B(batch 维 GEMM) : %8.2f ms %7.1f GFLOP/s 估算 DRAM 流量 %6.2f GB AI≈%.1f FLOP/byte\n",
sB * 1e3, flops / sB / 1e9, trB / 1e9, flops / trB);
printf("实测加速比 = %.1fx ;两版本 dW1 最大差异 = %.3e\n", sA / sB, err);
printf("每卡梯度缓冲大小 = %.2f MB(版本 A 每线程一份 → %d 线程共 %.1f MB)\n",
w1b / 1e6, omp_get_max_threads(), w1b * omp_get_max_threads() / 1e6);
return 0;
}
- 【代码做什么?】
- 数据准备:随机生成一批输入
X[B×D0]、标签Y[B]、两层权重W1[D1×D0]、W2[D2×D1],以及两套输出缓冲dW1a/dW1b,用同一份数据喂给两个版本,最后比较两者的dW1。 - 版本 A(
grads_per_sample):#pragma omp parallel开并行区,每个线程先分配自己的一整套梯度缓冲(lW1就是 D1×D0 = 160 万个 float = 6.42 MB),然后#pragma omp for schedule(static)把 B 个样本平均分给线程;每个样本内部串行完成”前向 → softmax 交叉熵梯度 → 两层外积累加”。所有线程跑完后,在一个critical段里把各自的私有缓冲加进全局梯度。 - 版本 B(
grads_batched):第一步先把整批的隐层激活 H 和上游梯度 G1/G2 全部算出来存下(这就是”激活内存”的来源);第二步把梯度计算重写成两个 GEMM:dW1 = G1ᵀ·X、dW2 = G2ᵀ·H。并行化方向从”样本”改成了”权重行“:每个线程独占若干行 dW1(无写冲突),行内以 batch 为累加轴——于是 X 的同一行会被 D1 个输出通道反复读取,且这些读取都命中缓存。 - 收尾:校验两个版本结果一致(只允许浮点误差),分别计时并打印 GFLOP/s、估算 DRAM 流量与算术强度。
- 数据准备:随机生成一批输入
- 【并行机制与性能解说】
- 线程/向量通道如何创建与分配:OpenMP 起 16 个线程,
schedule(static)把 B=256 个样本切成 16 块(每线程 16 个样本)。编译器的自动向量化会把内层for i < D0变成 AVX2/AVX-512 的vfmadd:以 AVX2(8 宽 FMA)为例,每个时钟节拍每核可完成 8 次乘加 = 16 FLOP;16 核 × 3.0 GHz × 8 宽 × 2 = 768 GFLOP/s 的峰值就是 4.1 节数值算例里用的那个数。 - 共享数据如何处理:
W1/W2/b1/b2/X/Y是只读共享(无需同步);dW1/dW2是写共享——版本 A 靠”每线程一份私有副本 + critical 合并”避免竞争,版本 B 靠”行所有权(row ownership)“避免竞争。注意版本 B 的行所有权顺带消灭了伪共享:每个线程独占整行(D0=784 个 float = 3136 字节 = 49 条 cache line),相邻线程不会写同一 cache line。反之,如果改成”每个线程负责若干列”,那么每次写dW1[j][i]都会和邻居共享 cache line → cache line ping-pong,性能可能掉一个数量级。 - Work / Span / 并行度:
- Work(一次迭代的浮点运算总数)≈
2·B·(D0·D1 + D1·D2)(前向)+4·B·(D0·D1 + D1·D2)(反向)≈ 6·B·D0·D1 = 6 × 256 × 784 × 2048 ≈ 2.47 GFLOP(D2=10 的贡献可忽略);两个版本 Work 完全相同。 - Span(关键路径):样本之间互相独立 → 关键路径就是”单个样本的前向+反向“的串行长度 ≈
6·D0·D1FLOP,再加上版本 A 末尾那次critical归约的串行时间(合并 6.42 MB 的全局梯度,这是 Span 里被硬生生加进来的一段)以及 OpenMP 的隐式 barrier。版本 B 的 Span = 单行 dW1 的累加链(B 次相关的外积累加)+ 两次并行 for 之间的 barrier。 - 并行度 = Work / Span:版本 A ≈ B(=256,样本数)——线程再多也没用,B 就是上限;版本 B ≈
D1 × (B/单行长度),此处 ≈ 2048 行 × 16 线程 = 32768 个独立任务,并行度远高于 16 个线程,因此 16 核能被喂满。结论:算法改写把可用并行度从 256 提到了 3 万,这才是版本 B 快的根本原因之一;另一个原因(更主要)是算术强度。
- Work(一次迭代的浮点运算总数)≈
- 瓶颈分析(这才是重点):
- 版本 A 是彻底的访存瓶颈。它每个样本都要把整块 dW1(6.42 MB)读一遍、写一遍。16 个线程各持一份私有副本 → 工作集 16 × 12.84 MB ≈ 205 MB,远超 L3,所以流量真实地落在 DRAM 上:
256 样本 × 19.27 MB ≈ 4.93 GB,算术强度只有2.47e9/4.93e9 ≈ 0.5 FLOP/byte(只看 dW1 的外积更新则是2 FLOP / 8 byte = 0.25 FLOP/byte)。按 20 GB/s 的单路内存带宽,模型预测耗时 ≈ 4.93 GB / 20 GB/s ≈ 247 ms——而这点计算量在 768 GFLOP/s 下只需 3.2 ms,差了近 77 倍。 - 版本 B 把算术强度抬了 400 倍:DRAM 流量降到 ≈ 11.5 MB(X 803 KB + H 2.1 MB + G1 2.1 MB + dW1 6.42 MB),
AI ≈ 214 FLOP/byte,于是耗时由计算决定:模型预测 ≈ max(2.47 GFLOP / 768 GFLOP/s, 11.5 MB / 20 GB/s) ≈ max(3.21 ms, 0.58 ms) ≈ 3.2 ms。两者的模型预测加速比 ≈ 77×。 - 这正是 CS149 讲义的核心结论在训练侧的复现:CS149 slide 9–10 用”循环融合把算术强度从 1/3 提到 3/5”来说明”提高算术强度让程序更可能变成计算受限“;slide 35 的分块 GEMM 与 slide 42–43 的 implicit GEMM 是同一招。训练里的”batch 维 GEMM”就是把 batch 融合进循环、让每个激活值被 D1 个权重复用——同一个优化,同一个收益。
- 另一个隐性代价:版本 B 要多存
H、G1、G2(此例 2.1+2.1+0.1 MB,很小;但对 Transformer 长序列会爆炸)。这就是”激活内存 vs 重算(recomputation / checkpointing)“的取舍,也是 2.4 节”内存账本未计入 activations”那句话的真正含义。
- 版本 A 是彻底的访存瓶颈。它每个样本都要把整块 dW1(6.42 MB)读一遍、写一遍。16 个线程各持一份私有副本 → 工作集 16 × 12.84 MB ≈ 205 MB,远超 L3,所以流量真实地落在 DRAM 上:
- 线程/向量通道如何创建与分配:OpenMP 起 16 个线程,
3.3 示例三:CUDA 分块梯度归约 —— 设备内的”层次化 AllReduce”
- 代码
// grad_partial.cu —— 权重梯度:按参数列并行 + split-K 部分和 + 块内树形归约
// 编译: nvcc -O3 -arch=sm_70 grad_partial.cu -o grad_partial
// 运行: ./grad_partial 4096 2048 4
// 参数: B(batch) D(特征维) C(batch 分块数 = split-K 度)
#include <cuda_runtime.h>
#include <cstdio>
#include <cstdlib>
#include <cmath>
#include <vector>
#include <algorithm>
#define CUDA_CHECK(x) do { cudaError_t e_ = (x); if (e_ != cudaSuccess) { \
printf("CUDA error %s at %s:%d\n", cudaGetErrorString(e_), __FILE__, __LINE__); exit(1);} } while (0)
// ---- 第一阶段:每个 block 负责一批参数列(threadIdx.x 连续 → 合并访存),
// 且只处理 batch 的一个分块(blockIdx.y = split-K 序号),写出部分和 ----
__global__ void grad_partial_kernel(const float* __restrict__ X, // [B x D]
const float* __restrict__ G, // [B x D] 上游梯度
float* __restrict__ part, // [C x D] 部分和
int B, int D, int chunk)
{
const int col = blockIdx.x * blockDim.x + threadIdx.x; // 参数列(= 权重列)
if (col >= D) return;
const int c = blockIdx.y; // batch 分块序号
const int n0 = c * chunk;
const int n1 = min(n0 + chunk, B);
float acc = 0.f;
for (int n = n0; n < n1; n++) // 串行累加 batch 维
acc += X[(size_t)n * D + col] * G[(size_t)n * D + col];
part[(size_t)c * D + col] = acc;
}
// ---- 第二阶段:一个 block 负责一列,把 C 个部分和在共享内存里做树形归约 ----
// 这就是"设备内版本的 allreduce 层次":块内 log2(blockDim) 步树形归约
__global__ void reduce_partial_kernel(const float* __restrict__ part,
float* __restrict__ dW, int D, int C)
{
extern __shared__ float s[];
const int col = blockIdx.x;
float acc = 0.f;
for (int c = threadIdx.x; c < C; c += blockDim.x) // 线程沿 C 维分工
acc += part[(size_t)c * D + col];
s[threadIdx.x] = acc;
__syncthreads();
for (int off = blockDim.x >> 1; off > 0; off >>= 1) { // log2(blockDim) 步树形归约
if (threadIdx.x < off) s[threadIdx.x] += s[threadIdx.x + off];
__syncthreads();
}
if (threadIdx.x == 0) dW[col] = s[0]; // 每列只写一次,无需原子操作
}
int main(int argc, char** argv)
{
const int B = (argc > 1) ? atoi(argv[1]) : 4096;
const int D = (argc > 2) ? atoi(argv[2]) : 2048;
const int C = (argc > 3) ? atoi(argv[3]) : 4;
const int chunk = (B + C - 1) / C;
std::vector<float> hX((size_t)B*D), hG((size_t)B*D);
for (size_t i = 0; i < hX.size(); i++) {
hX[i] = (float)((i * 2654435761u) % 1000) / 1000.f - 0.5f;
hG[i] = (float)((i * 40503u + 17) % 1000) / 1000.f - 0.5f;
}
float *dX, *dG, *dPart, *dW;
CUDA_CHECK(cudaMalloc(&dX, (size_t)B*D*sizeof(float)));
CUDA_CHECK(cudaMalloc(&dG, (size_t)B*D*sizeof(float)));
CUDA_CHECK(cudaMalloc(&dPart, (size_t)C*D*sizeof(float)));
CUDA_CHECK(cudaMalloc(&dW, (size_t)D*sizeof(float)));
CUDA_CHECK(cudaMemcpy(dX, hX.data(), (size_t)B*D*sizeof(float), cudaMemcpyHostToDevice));
CUDA_CHECK(cudaMemcpy(dG, hG.data(), (size_t)B*D*sizeof(float), cudaMemcpyHostToDevice));
const int TPB = 256;
dim3 gridP((D + TPB - 1) / TPB, C);
cudaEvent_t t0, t1;
cudaEventCreate(&t0); cudaEventCreate(&t1);
CUDA_CHECK(cudaDeviceSynchronize());
cudaEventRecord(t0);
for (int it = 0; it < 20; it++) {
grad_partial_kernel<<<gridP, TPB>>>(dX, dG, dPart, B, D, chunk);
reduce_partial_kernel<<<D, 128, 128 * sizeof(float)>>>(dPart, dW, D, C);
}
cudaEventRecord(t1);
CUDA_CHECK(cudaEventSynchronize(t1));
float ms = 0.f;
cudaEventElapsedTime(&ms, t0, t1);
ms /= 20.f;
double ref = 0.0; // CPU 参考:某一列的完整归约
for (int n = 0; n < B; n++) ref += (double)hX[(size_t)n*D + 7] * (double)hG[(size_t)n*D + 7];
float got = 0.f;
CUDA_CHECK(cudaMemcpy(&got, dW + 7, sizeof(float), cudaMemcpyDeviceToHost));
const double flops = 2.0 * B * (double)D;
const double bytes = 2.0 * B * (double)D * 4.0; // 读 X 与 G 各一次
printf("B=%d D=%d splitK=%d\n", B, D, C);
printf("时间 = %.4f ms ;算力 = %.1f GFLOP/s ;读入 %.2f MB/iter ;有效带宽 = %.1f GB/s\n",
ms, flops / (ms * 1e-3) / 1e9, bytes / 1e6, bytes / (ms * 1e-3) / 1e9);
printf("dW[7]: GPU = %.6f CPU = %.6f 误差 = %.3e\n", got, ref, fabs(got - ref));
printf("算术强度 = %.3f FLOP/byte(V100 峰值 15.7 TFLOP/s,ridge 点 ≈ 17.4 FLOP/byte)\n",
flops / bytes);
cudaFree(dX); cudaFree(dG); cudaFree(dPart); cudaFree(dW);
return 0;
}
- 【代码做什么?】
- 第一阶段(
grad_partial_kernel):grid 是二维的——blockIdx.x覆盖参数列(每 block 256 列,线程内连续 →X[n*D+col]与G[n*D+col]对同一个 n、相邻 col 的访问是完全合并的),blockIdx.y把 batch 切成 C 份(split-K)。每个线程串行累加自己那段 batch,写出一份部分和part[c][col]。 - 第二阶段(
reduce_partial_kernel):一个 block 负责一列,128 个线程沿 C 维分工(每个线程累加若干段),把结果放进共享内存,然后用for (off = 64, 32, ...)做 log₂(128) = 7 步树形归约,最后threadIdx.x == 0把结果写回dW[col]。 - 校验与计时:用 CPU 重算第 7 列的内积做参考;用 CUDA event 计时 20 次取平均,并打印算力、有效带宽与算术强度。
- 第一阶段(
- 【并行机制与性能解说】
- warp 与 SIMT 执行:256 线程的 block 分成 8 个 warp,每个 warp 的 32 个线程访问连续的 32 个 float = 128 字节,恰好是一次 128 B 的合并事务(4 个 sector)。
for (n = n0; n < n1; n++)是串行的、每次迭代一次 load,属于 high-latency low-ILP 的模式——靠同时驻留的大量 warp 做延迟隐藏(这就是 GPU 里 “occupancy 换延迟隐藏”的用法)。block 内没有线程间通信(除了第二阶段的树形归约),因此没有任何共享内存/原子争用。 - 共享数据如何处理:第一阶段完全无共享(每个线程写自己的一列的独立位置);第二阶段是层次化归约:线程内串行累加 → 块内共享内存树形归约 → 全局一次写入。这就是数据并行 allreduce 的微缩版:块内的树形归约 ≈ 节点内的 NVLink ring,块间的部分和 ≈ 节点间 allreduce,
splitK就是”用更多并发度换更长的累加链”。 - Work / Span / 并行度:
- Work = B×D 次乘加 = 2·B·D FLOP(此例 B=4096, D=2048 → 16.8 MFLOP);
- Span = 单列的累加链长度
nchunk = B/C = 1024次相关乘加 + 第二阶段的 7 步树形归约 + 一次同步; - 并行度 = Work/Span ≈
(B·D)/(B/C + log₂(128))≈ C·D = 8192 个独立累加任务 —— 远超 V100 的 80 个 SM × 2048 线程,并行度绝对充足,问题从来不是并行度。
- 瓶颈分析(本例最大的教学价值):
- 这个 kernel 的算术强度是
2·B·D / (2·B·D·4) = **0.25 FLOP/byte**——因为每个 X、G 的元素只被使用一次,没有任何复用。 - 在 V100 上(900 GB/s HBM、15.7 TFLOP/s FP32),它最多只能跑到
900 GB/s × 0.25 = **225 GFLOP/s**,即峰值的 **1.4%**。实测会非常接近这个带宽上限(67.1 MB / 900 GB/s ≈ 74.6 µs),CPU 上算这点活反而是浪费。 - 修法不是”加并行度”,而是”提高复用”:把梯度写成 GEMM(
dW = GᵀX),让每个 X 元素被 D 个权重列复用、每个 dW 元素在寄存器里累积,算术强度就从 0.25 升到D·B/(4B+2D) = 2048×4096/(16384+4096) ≈ **410 FLOP/byte**。这正是示例二版本 B 做的事,也是 cuDNN/CUTLASS 把卷积与注意力全部转成”分块 GEMM”的原因(CS149 slide 33–43:explicit GEMM 的 DRAM 流量会因 im2col 放大 R×S 倍,所以工业界改用 implicit GEMM——只把卷积矩阵的一个子块物化到共享内存,用分块 GEMM 例程吃掉它)。 - 另一个隐藏成本:
part[C×D]的写回与第二阶段的重读,让 X/G 之外又多了一趟显存往返(C×D×4 bytes ×2)。当 split-K 度 C 很大时(比如为了填满 SMs 而把 batch 切得很碎),这部分开销会显著增加——这是”并行度换通信量“的又一个实例。
- 这个 kernel 的算术强度是
- warp 与 SIMT 执行:256 线程的 block 分成 8 个 warp,每个 warp 的 32 个线程访问连续的 32 个 float = 128 字节,恰好是一次 128 B 的合并事务(4 个 sector)。
4. 性能模型与复杂度分析
4.1 训练迭代的 Work、Span 与 Roofline
一次数据并行迭代的时间可以写成三项之和:
\[T_{\text{iter}}(N) \;=\; \underbrace{\frac{W_{\text{compute}}}{N\cdot \Pi}}_{\text{算力分摊}} \;+\; \underbrace{T_{\text{comm}}(N)}_{\text{与 } N \text{ 近乎无关}} \;+\; \underbrace{T_{\text{sync}}(\alpha,\,N)}_{\text{步数} \times \text{启动延迟}}\]- Work:整个 iteration 的浮点运算量。对 Transformer 类模型有广泛使用的经验式 $W \approx 6PT$(P = 参数量,T = 本次迭代处理的 token 数;前向 2PT + 反向 4PT)。
- Span:关键路径 = “单层前向 → 单层反向 → 权重更新”的串行链,长度与层数 L 成正比,而与 N 无关。因此数据并行把 Work 分摊了,却没有缩短 Span——这正是它不像”算法并行”那样能无限扩展的原因。
- 并行度 = Work / Span ≈
6PT / (c·L)≈ P 量级(参数级)与 B 量级(样本级)。两个量都极大,所以数据并行的限制从来不是并行度,而是通信。 Roofline 视角(CS149 slide 6–8 的模型):设峰值算力 $\Pi$、内存带宽 $BW$,则可达吞吐 = $\min(\Pi, \; AI \times BW)$,拐点(ridge point)在 $AI^\* = \Pi/BW$。用 V100 的 FP16 数据:$AI^\* = 125\text{ TFLOP/s} / 900\text{ GB/s} \approx$ 139 FLOP/byte。
- 图 9:Roofline 上的几个关键工作点(横轴 log 刻度)
吞吐
(FLOP/s)
125 T ─┼────────────────────────────────────────────────┐ ← compute-bound 屋顶
│ ● d=12288,b=1024 线性层(≈877 F/B)
│ ╱
│ ╱ ← 屋顶线斜率 = 内存带宽 900 GB/s
│ ● 分块 GEMM(dW=GᵀX,≈410 F/B)
│ ╱
│ ╱
│ ● 梯度列并行 / 逐样本外积(0.25 F/B)→ 只有峰值的 1.4%
│ ╱ ● 融合前 1/3;● Ring AllReduce ≈ 0.25 F/B(与 N 无关)
└─┴────┴────┴────┴────┴────┴────┴────┴────┴────▶ 算术强度 FLOP/byte
0.25 1 2 4 8 16 32 64 139(ridge) 877
★ 所有"归约/通信"类工作都钉在图的左下角(AI ≤ 1),因此它们【永远】是带宽或延迟受限,
优化方向只能是【少搬字节】(fp16/bf16、梯度压缩、稀疏化、ZeRO 分片),而不是【少算 FLOP】。
表 4:本讲涉及的各阶段算术强度与受限类型(数值均为可复核的推导结果)
| 工作阶段 | 数据复用方式 | 算术强度(FLOP/byte) | V100 上的受限类型 | 优化手段 |
|---|---|---|---|---|
| 逐样本权重外积累加(示例二 A、示例三) | 无复用,每次读改写整块 dW | 0.25 | 严重带宽受限(~1.4% 峰值) | 改成 batch 维 GEMM |
| AllReduce 梯度(任意 N,fp16) | 每个元素只被加一次 | 0.25(与 N 无关) | 带宽/延迟受限 | 减字节、减步数、换拓扑 |
| 分块 GEMM 算 dW(示例二 B) | X 被 D1 行复用 | ~214(本例) | 计算受限 | 加 cache 分块、提 tile 尺寸 |
| 大 batch 线性层(d=12288, b=1024) | 权重被整批复用 | ~877 | 计算受限(Tensor Core) | 提 batch、混合精度 |
| 大 batch 线性层(d=12288, b=1) | 只复用一次 | ~1 | 权重带宽受限 | 攒 batch / 权重驻留 |
线性层 AI 的推导(可直接代数字):
AI = 2·b·d² / (2·d² + 4·b·d) = b·d/(d + 2b)。代入 b=1024, d=12288 得 877;代入 b=1 得 ≈1.0。这是”batch 越大越划算”的定量解释,也是数据并行本身能存在的前提——正是大 batch 把训练推进了计算受限区。
4.2 强扩展(strong scaling)与弱扩展(weak scaling)
数据并行的两种扩展方式在时间模型上差别巨大,值得单独列出:
表 5:强扩展 vs 弱扩展($T_c$ = 单卡跑完整 batch 的计算时间,$T_{\text{comm}}$ = 一次 AllReduce 时间)
| 维度 | 强扩展(固定全局 batch) | 弱扩展(固定每卡 batch) |
|---|---|---|
| 每卡计算量 | $T_c/N$(随卡数下降) | $T_c$(恒定) |
| 通信量 | $T_{\text{comm}}$(恒定,与 N 无关) | $T_{\text{comm}}$(恒定) |
| 迭代时间 | $T_c/N + T_{\text{comm}}$ | $T_c + T_{\text{comm}}$(恒定) |
| 加速比 $S(N)$ | $1/(1/N + r)$,其中 $r = T_{\text{comm}}/T_c$ | $N$(吞吐线性增长,效率恒定) |
| 加速比上限 | $\mathbf{1/r = T_c/T_{\text{comm}}}$ | 无通信上限;但受收敛性限制(batch 太大时每步收益递减) |
| 失效模式 | 卡加到某个数量后完全无收益 | 需要同步放大学习率;过大 batch 会浪费算力 |
| 适用场景 | 想更快得到结果(时间受限) | 想训得更大(问题规模受限) |
注意区分两种”Amdahl”:经典 Amdahl 律管的是串行部分(这部分占比固定);数据并行里的通信是常数开销(不随 N 缩小),得到的规律是”加速比天花板 = 计算时间/常数开销”。两者形式相似,来源不同,不要混用。
4.3 数值算例一:GPT-2 1.5B 在 8 张 V100 上的数据并行
假设(全部显式给出,便于复核):M = 1.5×10⁹ 参数(对应讲义表里的 GPT-2 量级);混合精度梯度 = 2 B/参数;每卡 batch = 1024 个 token,全局 batch = 8×1024 = 8192 token;每迭代算力 $W = 6PT = 6 \times 1.5\times10^9 \times 8192 = 7.37\times10^{13}$ FLOP = 73.7 TFLOP;单卡实测算力按峰值 125 TFLOP/s 的 40% MFU 取 50 TFLOP/s;环上 AllReduce。
- 梯度缓冲:$1.5\times10^9 \times 2\,\text{B} = $ 3.0 GB。
- 每卡 Ring AllReduce 通信量:$2(N-1)/N \times 3.0\ \text{GB} = 2\times7/8\times3.0 =$ 5.25 GB(≈ 2M 个参数,与 N 无关)。
- 每卡计算时间:$6PT/N = 73.7/8 = 9.22$ TFLOP,除以 50 TFLOP/s = 184 ms。
- 通信时间(不同互连):
表 6:同一个 1.5B 模型、同一个 8 卡作业,换一条链路就差一个数量级
| 互连 | 单向有效带宽 | AllReduce 时间 | 迭代总时间(184 ms + 通信) | 相对单卡加速比 | 加速比天花板 $1/r$ |
|---|---|---|---|---|---|
| NVLink 2.0(卡间) | 150 GB/s | 35 ms | 219 ms | 6.70×(84% 效率) | 42× |
| 400 Gb/s InfiniBand | 50 GB/s | 105 ms | 289 ms | 5.08× | 14× |
| PCIe Gen3 x16 | 12.6 GB/s | 417 ms | 601 ms | 2.45×(31% 效率) | 3.5× |
| 100 Gb/s Ethernet | 12.5 GB/s | 420 ms | 604 ms | 2.43× | 3.5× |
| 40 Gb/s Ethernet | 5 GB/s | 1050 ms | 1234 ms | 1.19×(15% 效率) | 1.4× |
读数:单卡跑完整 batch 需要 73.7 TFLOP / 50 TFLOP/s = 1.47 s。用 8 张 V100 走 NVLink,加速 6.7×;走 PCIe 或 100 GbE,只有 2.45×——8 张卡里有一半的算力被浪费在等梯度。走 40 GbE 时更是只快 1.19×,几乎等于没并行。天花板 $1/r = T_c/T_{\text{comm}}$ 告诉你”这台机器最多值几张卡”:PCIe 集群上超过 4 张卡就没有意义了。这也解释了为什么真实的分布式训练中心一定用 NVLink + InfiniBand,并把通信与反向计算重叠(ZeRO-2 的”每层反完立刻 Reduce”就是这个目的)。
4.4 数值算例二:AllReduce 四种算法的 cross-over(含可运行模型程序)
- 代码
/* allreduce_model.c —— Ring / Tree / Butterfly / Parameter Server 的通信代价模型
* 编译: gcc -O2 -std=c11 allreduce_model.c -o allreduce_model -lm
* 运行: ./allreduce_model 1500000000 2 5e-6 1e-10
* 参数: M(参数量) 每参数字节数 alpha(每步启动延迟,秒) 每字节传输时间(秒) → 10 GB/s
*/
#include <stdio.h>
#include <stdlib.h>
#include <math.h>
static double naive(double M, int N, double w, double a) { return (N - 1) * M * w + a; }
static double ps(double M, int N, double w, double a) { return N * M * w + 2 * a; }
static double ring(double M, int N, double w, double a) { return 2.0 * (N - 1) * (a + (M / N) * w); }
static double tree(double M, int N, double w, double a) { return 2.0 * ceil(log2((double)N)) * (a + M * w); }
static double bfly(double M, int N, double w, double a) { return ceil(log2((double)N)) * (a + M * w); }
int main(int argc, char **argv)
{
const double M = (argc > 1) ? atof(argv[1]) : 1.5e9; /* 参数个数 */
const double be = (argc > 2) ? atof(argv[2]) : 2.0; /* 每参数字节数(fp16 = 2) */
const double a = (argc > 3) ? atof(argv[3]) : 5e-6; /* 每步启动延迟(秒) */
const double bw = (argc > 4) ? atof(argv[4]) : 1e-10; /* 每字节时间 → 10 GB/s */
const double w = bw / be; /* 每个参数元素的传输时间 */
printf("M = %.3g 个参数(%.2f GB,每参数 %g B);alpha = %g us;链路 = %.1f GB/s\n\n",
M, M * be / 1e9, be, a * 1e6, 1.0 / bw / 1e9);
printf("%6s %12s %12s %12s %12s %12s\n",
"N", "PS(ms)", "Naive(ms)", "Ring(ms)", "Tree(ms)", "Butterfly(ms)");
for (int N = 2; N <= 1024; N *= 2)
printf("%6d %12.2f %12.2f %12.2f %12.2f %12.2f\n", N,
ps(M,N,w,a)*1e3, naive(M,N,w,a)*1e3, ring(M,N,w,a)*1e3,
tree(M,N,w,a)*1e3, bfly(M,N,w,a)*1e3);
return 0;
}
【代码做什么?】 把表 1 里的五个公式直接实现成函数:
ring用 $2(N-1)(\alpha + (M/N)w)$(对应讲义”每步送 M/N,重复 2N 次”),tree用 $2\log_2 N(\alpha + Mw)$(每步送整份 M),bfly用 $\log_2 N(\alpha + Mw)$,ps用 $N M w$(服务器串行服务 N 个 worker),naive用 $(N-1)Mw$。程序扫过 N = 2…1024,打印一张延迟表。【并行机制与性能解说】 这个程序本身串行,它的价值是把讲义 slide 25–26 的对比表变成可以自己调参的实验:调
be(fp16↔fp32)、调bw(NVLink↔以太网)、调M(小模型↔大模型),就能看到”哪种 AllReduce 更好”取决于工作点。Work/Span 视角:五种算法的 Work 完全相同(都是把 M 个参数加起来),差别全在 Span(步数与每步大小)——这正是”同一份工作,不同的关键路径”的教科书例子。数值算例一:大模型(M = 1.5×10⁹,fp16 = 3.0 GB,链路 10 GB/s,α = 5 µs)
单次传输整份梯度的带宽成本 $M w = 1.5\times10^9 \times 5\times10^{-11} = 75$ ms。
表 7:AllReduce 五种方案的延迟(ms)—— 大模型、带宽主导区间
| N | Parameter Server | Naïve | Ring | Tree | Butterfly |
|---|---|---|---|---|---|
| 2 | 150.01 | 75.00 | 75.01 | 150.01 | 75.00 |
| 8 | 600.01 | 525.00 | 131.32 | 450.03 | 225.02 |
| 64 | 4800.01 | 4725.00 | 148.29 | 900.06 | 450.03 |
| 1024 | 76800.01 | 76725.00 | 160.08 | 1500.10 | 750.05 |
读数:M 很大时 Ring 完胜,而且从 N=8 到 N=1024 只从 131 ms 涨到 160 ms(趋近 2Mw = 150 ms 的平台,增量全部来自 $2(N-1)\alpha$ 的启动延迟)。这正面回答了讲义 slide 25 的问题:Ring 的单卡通信量 ≈ 2M 与 N 无关,所以 N 增大时带宽成本不变、只有微小的步数延迟增长;而 PS 与 Naive 的成本正比于 N(N=1024 时 76.8 s,比 ring 慢 480 倍);Tree/Butterfly 因为”每步要传整份 M”,成本随 $\log N$ 增长,在这里反而吃亏。
数值算例二:小模型(M = 10⁶,fp16 = 2 MB,同样 10 GB/s,α = 5 µs)
此时 $M w = 50$ µs,单步传输时间与启动延迟同量级,结论完全反转:
表 8:同样五种方案的延迟(ms)—— 小模型、延迟主导区间
| N | Parameter Server | Naïve | Ring | Tree | Butterfly |
|---|---|---|---|---|---|
| 8 | 0.41 | 0.35 | 0.16 | 0.33 | 0.17 |
| 64 | 3.21 | 3.15 | 0.73 | 0.66 | 0.33 |
| 1024 | 51.21 | 51.15 | 10.33 | 1.10 | 0.55 |
读数:M 小的时候 Butterfly / Tree 反超 Ring——N=1024 时 Butterfly 只用 0.55 ms(log₂1024 = 10 步 × 55 µs),而 Ring 要 10.33 ms(2046 步 × 5 µs,即 10 ms 全部是启动延迟)。这就是 cross-over:带宽主导时选 Ring(少传字节),延迟主导时选 Tree/Butterfly(少走轮次)。真实的集合通信库(NCCL)因此采用混合拓扑:节点内用 NVLink 做 ring,节点间用树/多环,再按消息大小在算法之间切换。
4.5 数值算例三:ZeRO 到底省了多少、又贵了多少
假设:M = 175×10⁹(GPT-3 量级),每参数 20 B(混合精度训练账本),N = 64 张卡,机间 400 Gb/s InfiniBand(50 GB/s),单卡 FP16 实测 125 TFLOP/s(A100 的 40%),本次迭代处理 T = 3.2×10⁶ 个 token(大 batch)。
- 显存(代入表 3):
- 基线数据并行:3500 GB/卡 → 完全不可能;
- ZeRO-1:700 + 2800/64 = 743.8 GB/卡 → 不可能;
- ZeRO-2:350 + 18×175/64 ≈ 399.2 GB/卡 → 不可能;
- ZeRO-3:3500/64 = 54.7 GB/卡 → A100-80GB 装得下(还要留激活的内存)。 结论:要把 175B 喂进 8×80 GB 的机器,必须上 ZeRO-3(或叠加张量并行/流水并行)。
- 通信量:每卡每迭代 ≈ 2M 个参数(基线/ZeRO-1/2)= 700 GB,ZeRO-3 = 3M = 1050 GB。
- 通信时间:700 GB ÷ 50 GB/s = 14 s(不重叠时)。
- 计算时间:$6PT/(N\cdot\Pi) = 6\times175\times10^9\times3.2\times10^6 / (64 \times 125\times10^{12}) = 3.36\times10^{18}/8\times10^{15} =$ 420 s。
- 对比:14 s / 420 s = 3.3%——大 batch 下通信占比很小,训练是计算受限的。但若把 batch 缩小到 8192 个 token(T = 8192):计算时间 = $6\times175\times10^9\times8192/8\times10^{15} =$ 1.07 s,通信 14 s 反过来是主导(通信占比 93%)。
- cross-over batch size:令 $6PT/(N\Pi) = V/BW$ 解得 \(T^\* = \frac{V\cdot N\cdot\Pi}{BW\cdot 6P} = \frac{700\times10^9 \times 64 \times 125\times10^{12}}{50\times10^9 \times 6\times 175\times10^9} \approx 1.07\times10^5\ \text{token}\) 即 每次迭代的全局 batch 超过约 10.7 万 token 之后,通信才不再是瓶颈。 这解释了分布式训练中所有”增大 batch + 同步放大学习率”的工程实践:大 batch 不是为了更快收敛,而是为了把通信在计算后面藏起来。(不过 batch 不能无限放大:公开文献中普遍观察到超过某个”临界 batch size”后每步的收敛增益会急剧衰减——这一点讲义没有展开,属公开领域的补充知识。)
4.6 把开销放进同一张预算表
把 4.3 节的算例拆成”随 N 下降”与”不随 N 下降”两类,就能一眼看出优化该往哪里使劲:前向 92 ms、反向 92 ms 属于∝1/N 的可分摊项(靠混合精度、算子融合、激活重算继续压缩);AllReduce 梯度 35 ms(NVLink)/ 417 ms(PCIe)与 2(N−1)α 的同步延迟属于恒定项(只能靠 fp16/bf16、梯度压缩、ZeRO-2 分桶重叠、合并小消息来压缩);Adam 更新在单卡上数值极小可忽略,但它是 ZeRO-1 分片的直接受益者。加速比天花板的本质就是”恒定项 / 可分摊项”:把 417 ms 压到 35 ms,天花板就从 3.5× 抬到 42×。
5. 关键要点
数据并行的可并行性来自 SGD 求和号里的”样本维”,代价是”参数维必须整体同步”。 前向/反向的计算量随卡数 N 线性分摊,而梯度聚合的通信量 ≈ 2M 个参数、与 N 几乎无关——于是加速比存在天花板 $1/r = T_c/T_{\text{comm}}$(本讲算例:NVLink 42×、PCIe 仅 3.5×)。判断一个集群值不值得加卡,先算这个天花板。
AllReduce 的四种算法是”通信量 vs 轮数”的连续权衡,没有全局最优。 Ring 用”每卡只搬 2M(与 N 无关)+ 2(N−1) 步”取得最佳可扩展性(N=1024 时仍只需 2Mw);Tree/Butterfly 把轮数压到 $2\log N$ / $\log N$,代价是每步搬整份 M,因此在小模型 + 大 N 时才反超(N=1024、M=10⁶ 时 Butterfly 0.55 ms vs Ring 10.33 ms)。选算法要先判断工作点落在带宽主导区还是延迟主导区。
“去中心化”是扩展性的前提。 参数服务器的中心带宽随 N 线性劣化($MN/BW$),Naive AllReduce 是 O(N²) 通信量——两者在 N=1024 时比 Ring 慢两个数量级以上。任何”所有节点都去访问同一个东西”的设计都会在规模化时崩掉,这是并行系统设计的通用准则。
大模型训练的瓶颈首先不是算力而是”内存账本”,其次是”传输字节数”。 混合精度训练要 20 B/参数(FP16 参数 + FP16 梯度 + 16 B FP32 优化器状态),而且不含激活;ZeRO 通过把这份账本切成 1/N 换来显存,其中 Stage 1/2 的通信量与基线完全相同(免费省内存),只有 Stage 3 需要多付 50%。所有通信优化的正确方向是”少搬字节“(fp16/bf16、压缩、分片),而不是”少算 FLOP”——因为归约类工作的算术强度只有 0.25 FLOP/byte 且与 N 无关(每个元素搬进搬出共 4 字节、只摊到 1 次加法),在 Roofline 上永远钉在左下角。
单卡内部和集群之间的优化是同一套道理的两个尺度。 单卡上要用分块 GEMM 把算术强度从 0.25 FLOP/byte(逐样本外积)抬到 10² 以上(batch 维 GEMM、implicit GEMM、Flash-Attention),否则算力利用率不到 2%;集群上则要用好的归约拓扑把通信量压到 2M 并让它与计算重叠(分桶 AllReduce)。“提高算术强度”和”减少通信字节”是贯穿全课程的两条主线。
6. 常见陷阱与注意事项
误以为”卡越多越快”。 通信量 ≈ 2M 与 N 无关,而计算量 ∝ 1/N,因此 $S(N) = 1/(1/N + r)$ 会迅速饱和(PCIe 上 8 卡的效率只有 31%,40 GbE 上只有 15%)。先用
T_c/T_comm算天花板,再决定买几张卡;加了卡但 batch 不放大时,很可能是在给网卡打工。把”AllReduce 时间”当成与 N 无关的常数。 只有带宽项 $2Mw$ 与 N 无关;延迟项 $2(N-1)\alpha$ 随 N 线性增长。小模型、多节点时这一项会主导(示例:M = 10⁶、N = 1024 时 10.33 ms 几乎全是启动延迟)。小消息要合并(bucketing),否则每次迭代都在给网络交”起步价”。
在梯度归约里制造数据竞争或伪共享。 数据并行最容易犯的错是”多个线程/block 直接
+=同一块 dW”:正确做法是(a)层次化归约(线程私有 → 块内树形归约 → 每块一次原子操作,如示例三)、或(b)行所有权(每线程独占若干权重行,如示例二版本 B)。若按列切分 dW,相邻线程会写同一 cache line → 伪共享(false sharing),性能可掉一个数量级。真实框架里梯度累加用atomicAdd/__shfl归约,正是这个原因。忘记”梯度缓冲是每卡一份全量”这件事会随卡数放大显存。 数据并行不省显存:N 张卡存 N 份完全相同的副本(1.5B 模型的 fp16 梯度就是 3 GB/卡,另加 20 B/参数的训练状态)。显存不够时不要只想着”减 batch”,要先算清账本,再决定用 ZeRO 哪一级。
用逐样本外积的方式算权重梯度。 这是最典型的”算得动但跑不快”:算术强度 0.25 FLOP/byte,在 V100 上只能跑到峰值的约 1.4%(示例三实测会贴近 900 GB/s 带宽上限)。正确做法是把 batch 融合进循环,用分块 GEMM 让每个激活值被 D 个输出通道复用(示例二版本 B,AI 从 0.5 提到 214 FLOP/byte,模型预测加速 77×)。
假设浮点归约的结果可复现、或忽略激活内存。 Ring AllReduce 的累加顺序随 N 改变,换卡数就会改变数值结果的最低有效位,测试必须用容差;同时”内存账本”通常不含输入 batch 与全部激活(讲义 slide 40 明确注明),长序列 Transformer 的激活可以比参数本身还大(示例:32 层 × 16 头、序列 8192 时,朴素注意力要物化 8192² 的矩阵,合计约 68.7 GB;分块/Flash 式融合只需约 1.07 GB),必须靠激活重算或分块融合来控制。
7. 思考题(带答案)
问题 1(拓扑选择):某团队要在 64 个节点(每节点 8 张卡)上训练一个 参数量只有 4×10⁶ 的小模型,机间网络是 100 Gb/s Ethernet(单向 12.5 GB/s),每步的启动延迟 α = 5 µs。训练代码目前用 Ring AllReduce 同步梯度(fp32,每参数 4 B)。请估算一次 AllReduce 的耗时,指出瓶颈在哪一项,并给出至少两条改进方案及其定量收益。
【答案】 先算基本量:M = 4×10⁶ 参数,fp32 即 16 MB。$w = 1/(12.5\times10^9) = 8\times10^{-11}$ s/byte,故每参数传输时间 = 4 B × 8×10⁻¹¹ = 3.2×10⁻¹⁰ s,整份 M 的带宽成本 $Mw = 4\times10^6 \times 3.2\times10^{-10} = 1.28$ ms。N = 64×8 = 512 卡。
- Ring:$2(N-1)(\alpha + (M/N)w) = 2\times511\times(5\ \mu s + 1.28\ \text{ms}/512) = 1022 \times (5 + 2.5)\ \mu s = 1022 \times 7.5\ \mu s \approx$ 7.7 ms。其中 $1022\times5\ \mu s = 5.11$ ms 是纯启动延迟(占 66%),带宽只有 2.6 ms。瓶颈是”步数 × α”,即延迟项,而不是带宽。
- 对照 Butterfly:$\log_2 512 = 9$ 步 × (5 µs + 1.28 ms) = 9 × 1.285 ms ≈ 11.6 ms——反而更慢!因为每步要传整份 M,带宽成本放大 9 倍。这正是”小模型不一定该用 butterfly”的定量证据:本例中 per-step 带宽成本(1.28 ms)远大于 α(5 µs),所以”少传字节”比”少走轮次”更值钱。
- Tree:$2\log_2 512 \times 1.285\ \text{ms} = 18 \times 1.285 =$ 23.1 ms,最差。
- 改进方案(每条都给出定量收益):
- 改成 fp16/bf16 传梯度(2 B/参数):每参数时间减半,带宽成本从 2.6 ms 降到 1.3 ms,Ring 总时间 7.7 ms → 6.4 ms(再配合 loss scaling 保证收敛)。
- 分层归约(节点内 NVLink ring + 节点间 ring):节点内 8 卡走 NVLink(150 GB/s),每节点只由 1 张卡把”本节点已归约的 16 MB”发到机间。机间步数从 1022 降到 2×(64−1) = 126 步,机间延时项 126×5 µs = 0.63 ms,机间带宽项 2×(16 MB/12.5 GB/s) = 2.6 ms,节点内几乎可忽略 ⇒ 总计约 3.3 ms(相对 7.7 ms 提升 2.3×)。
- 梯度分桶 + 与反向计算重叠:把 16 MB 梯度切成若干 bucket,每算完一层的梯度就发起该 bucket 的 AllReduce,让 5.11 ms 的启动/传输延迟藏在反向计算后面;理论上通信可被完全隐藏,端到端开销趋近于 0(前提是反向计算时间 ≥ 通信时间)。
- 增大 batch(弱扩展):每卡 batch 增大不会改变通信时间,但会让”计算时间/通信时间”之比上升,从而让通信占比下降——这是最省事的一条,代价是收敛性需要重新调参。
问题 2(内存与 ZeRO 的取舍):某团队在 8 张 A100-80GB 上训练一个 17.2B 参数的模型(与讲义中 Turing NLG 同量级),采用混合精度(20 B/参数的训练账本,另有每卡约 12 GB 的激活)。请判断以下三种配置是否可行,并说明各自的通信代价;如果都不可行,指出最少需要多少张卡才能让 ZeRO-3 跑起来。
【答案】 训练状态总量 = 17.2×10⁹ × 20 B = 344 GB(这就是表 3 里”基线每卡 344 GB”的来源;讲义表中给出的 275 GB 是 16 B/参数的口径,不含 FP16 参数与 FP16 梯度的额外 4 B)。
- 基线数据并行:每卡仍要 344 GB 全量 → 344 + 12 = 356 GB/卡 > 80 GB,不可行。
- ZeRO-1(每卡 $4M + 16M/N$):$4\times17.2 = 68.8$ GB + $16\times17.2/8 = 34.4$ GB = 103.2 GB,再加 12 GB 激活 = 115.2 GB/卡 > 80 GB,不可行。通信量 = 2M 元素(与基线相同)。
- ZeRO-2(每卡 $2M + 18M/N$):$34.4 + 18\times17.2/8 = 34.4 + 38.7 = 73.1$ GB,加 12 GB 激活 = 85.1 GB/卡 > 80 GB,仍不可行(差得很近)。
- ZeRO-3(每卡 $20M/N$):$344/8 = 43$ GB + 12 GB 激活 = 55 GB/卡 < 80 GB,可行;通信量 = 3M 元素 = 基线的 1.5 倍。
- 最少卡数:要求 $344/N + 12 \le 80 \Rightarrow N \ge 344/68 \approx 5.06$,即至少 6 张 A100-80GB(留出碎片与临时缓冲的现实余量后,8 张是合理配置)。若叠加张量并行把激活与参数进一步切开,还可以更省——这正是讲义 slide 66 说 “Turing NLG 17.2B is powered by Stage 1 and Megatron” 的含义:真实的大模型训练从来不是单一并行方式,而是 ZeRO-1 + 张量并行的组合。
- 补充要点:ZeRO-3 的 1.5 倍通信在小 batch 时会成为主要瓶颈(第 4.5 节的 cross-over 分析:batch 不足约 10 万 token 时通信占主导),因此实践中要配合”参数 AllGather 与计算流水重叠”和非连续的梯度 bucket 划分。若显存够用,优先选 ZeRO-2 而不是 ZeRO-3:省下 50% 通信。
问题 3(性能模型与代码判断):下面这段代码在一个 16 核 CPU(峰值 768 GFLOP/s,内存带宽 20 GB/s)上训练一个小 MLP。请指出它慢的根本原因,给出定量分析,并写出改进后的循环结构。
// B = 256 个样本,D0 = 784,D1 = 2048
for (int n = 0; n < B; n++) {
// ... 前向得到第 n 个样本的隐层梯度 g[D1] ...
for (int j = 0; j < D1; j++)
for (int i = 0; i < D0; i++)
dW1[j][i] += g[j] * X[n][i]; // 逐样本外积累加
}
【答案】
- 慢的根本原因:算术强度极低,是彻底的访存瓶颈,而不是算力不足。 这个双层循环对
dW1(D1×D0 = 160 万个 float = 6.42 MB)每处理一个样本就读一遍、写一遍,只在寄存器里做一次乘加就丢掉。每个 X 元素、每个 dW1 元素的复用次数都是 1。 - 定量分析:每样本的乘加次数 = D0×D1 = 1.605×10⁶ 次乘法(2 FLOP)⇒ 3.21 MFLOP;访存 = 读 dW1 6.42 MB + 写 dW1 6.42 MB = 12.84 MB ⇒ AI = 3.21×10⁶/12.84×10⁶ = 0.25 FLOP/byte。整批 256 个样本:计算 822 MFLOP,访存 3.29 GB。按 20 GB/s 带宽,至少 165 ms;而按 768 GFLOP/s 算力只要 1.07 ms —— 访存时间是计算时间的 154 倍。可用带宽上限 = 20 GB/s × 0.25 = 5 GFLOP/s,即峰值的 0.65%。
- 改进后的循环结构(batch 维 GEMM + 行所有权):
// 先算出整批的上游梯度 G1[B][D1] 与激活,再一次性做 dW1 = G1^T · X
#pragma omp parallel for schedule(static)
for (int j = 0; j < D1; j++) { // 并行在权重行上:每线程独占若干行(无竞争、无伪共享)
float *row = dW1[j]; // 或 &dW1[j*D0];整行常驻 cache 反复复用
for (int n = 0; n < B; n++) { // batch 维作为累加轴 → 复用率 = D1
const float g = G1[n][j];
const float *x = X[n];
for (int i = 0; i < D0; i++) row[i] += g * x[i];
}
}
- 收益:dW1 只在最后写回一次,X 的每一行被 D1 = 2048 个输出通道复用(且命中缓存),DRAM 流量降到约
X(803 KB) + G1(2.1 MB) + dW1(6.42 MB) ≈ 11.5 MB,AI 提升到822×10⁶/11.5×10⁶ ≈ 214 FLOP/byte。此时耗时由计算决定:max(822 MFLOP/768 GFLOP/s, 11.5 MB/20 GB/s) = max(1.07 ms, 0.58 ms) ≈ 1.07 ms,相对原来的 165 ms 约 150 倍(示例二中因为还要存取 W1,实测的模型预测加速比是 77×)。 - 同时要说清的代价:这个改写必须把整批的中间激活与上游梯度留在内存里(激活内存随 batch 线性增长),这正是”用内存换算术强度”的典型交换,也是真实框架里用激活检查点(checkpointing / recomputation)在两者之间找平衡的原因。此外,若把并行方向错选成”按列切分 dW1”(每线程负责若干列 i),相邻线程会写同一 cache line,将引入伪共享,收益会被吃掉一大截——必须按行(连续维度)切分。
