LLM WIKI · 课程精读

LEARNING UNIT · 14

GEMM、Tensor Core 与分布式并行

追踪矩阵乘的数据复用,并连接双层缓冲、Tensor Core、张量并行和通信原语。

已整理章节
11 节
单元来源
10 条视频
总时长
31:45
状态
已发布
学习位置
14 / 20
01

主题讲解 · 03:51

从标量内积到分块:读懂 GEMM 的访存优化

学习目标

  • 写出包含 α\alphaβ\beta 的 GEMM 完整定义,并区分输入矩阵与输出矩阵。
  • CijC_{ij} 的行列点积解释朴素矩阵乘法。
  • 说明为什么只数浮点运算不能判断 GPU kernel 是否高效。
  • 画出一个输出 tile 对应的 AABB tile 访问路径。
  • 解释 shared memory 与寄存器在分块 GEMM 中承担的不同角色。
  • 在视频的理想模型下推导 tile 大小怎样降低 HBM 访存量。

前置与衔接

前面的课程已经讨论了张量计算、算子融合和 GPU 存储层次。

本课把这些概念集中到最典型的算子 GEMM:

Cout=αAB+βCin.C_{\text{out}}=\alpha AB+\beta C_{\text{in}}.

矩阵乘本身并不神秘,难点是让片上已加载的数据被尽可能多次复用。

理解本课后,下一课的 tensor parallel 才能同时从“矩阵 shape”和“通信位置”两条线阅读。

核心讲解

1. GEMM 比单纯的矩阵乘多了两个缩放项

通用 GEMM 接口计算

Cout=αAB+βCin.C_{\text{out}}=\alpha AB+\beta C_{\text{in}}.

其中 α\alpha 缩放乘积,β\beta 决定旧的 CinC_{\text{in}} 是否参与累加。

图 1

GEMM 的一般形式是 C_out=αAB+βC_in;深度学习中的纯矩阵乘常取 α=1、β=0。

原视频 · 00:20 ↗

视频为了突出矩阵乘,把深度学习中常见的简单情形写成 α=1,β=0\alpha=1,\beta=0

这不是 GEMM 的普遍固定参数;残差累加、线性组合或某些 BLAS 调用会使用其他取值。

为避免输入与输出都叫 CC 造成混淆,这里显式记为 CinC_{\text{in}}CoutC_{\text{out}}

2. 一个输出元素是一行与一列的点积

ARM×K,BRK×N,A\in\mathbb R^{M\times K},\qquad B\in\mathbb R^{K\times N},

则乘积 C=ABRM×NC=AB\in\mathbb R^{M\times N},且

Cij=k=0K1AikBkj.C_{ij}=\sum_{k=0}^{K-1}A_{ik}B_{kj}.
图 2

输出元素 C_ij 等于 A 的第 i 行与 B 的第 j 列在 K 维上的内积。

原视频 · 00:40 ↗

朴素实现可以为每个 CijC_{ij} 独立读取一整行 AA 和一整列 BB

算术是正确的,但相邻输出元素会反复读取许多相同数据。

3. 访存瓶颈来自“用一次就丢”

视频用 N×NN\times N 方阵简化分析。

若完全忽略缓存与复用,每个输出元素要读取 NNAA 元素和 NNBB 元素。

全体 N2N^2 个输出元素的输入读取量因此按

O(N2N)=O(N3)O(N^2\cdot N)=O(N^3)

增长。

这里的 O(N3)O(N^3) 是视频采用的“最坏朴素 HBM 读取”模型,不是说所有 GEMM 实现都真的发出这么多显存事务。

硬件 cache、coalescing 和库实现都会改变常数乃至实际流量。

4. 分块先确定一个要完成的输出 tile

视频以 N=128N=128、tile 边长 T=32T=32 为例。

输出矩阵被划分为 4×44\times432×3232\times32 tile。

图 3

把 128×128 矩阵按 TILE=32 切成 4×4 小块,一个红色 C tile 由对应 A tile 行与 B tile 列累加得到。

原视频 · 01:40 ↗

设目标输出块为 CI,JC_{I,J}

沿收缩维 KK,它需要依次计算

CI,Jmathrel+=AI,kBk,J,C_{I,J}mathrel{+}=A_{I,k}B_{k,J},

其中每个 AI,kA_{I,k}Bk,JB_{k,J} 都是 T×TT\times T 小矩阵。

5. HBM 负责供给,shared memory 负责块级复用

每轮把一对 AABB tile 从 HBM 搬到片上存储,再由线程块内的多个线程复用。

图 4

A、B 的对应 tile 分批从 HBM 搬到片上 SRAM,有限空间中用新 tile 替换旧 tile。

原视频 · 02:20 ↗

在常见 CUDA 叙述中,这块由程序员协作管理的片上空间是 shared memory。

shared memory 的物理实现位于片上 SRAM,但不能反过来说“GPU 所有片上 SRAM 都是 shared memory”。

cache、寄存器文件等也属于片上资源,却有不同接口与作用域。

6. 寄存器保存输出累加器

线程从 shared memory 读取小片段,执行乘加,并让所负责的输出元素长期停留在寄存器中。

图 5

各组 tile 的乘积持续累加到寄存器中的同一个 C tile,完成前不反复写回 HBM。

原视频 · 02:40 ↗

沿 KK 维的 tile 全部处理完以后,完成的 CC tile 才写回 HBM。

因此层次化复用可以概括为:

  • HBM 到 shared memory:同一个输入 tile 被线程块内多个计算复用;
  • shared memory 到寄存器:片段被更细粒度的乘加复用;
  • 寄存器到 HBM:输出部分和避免每次乘加都回写显存。

7. 理想模型下,tile 边长带来约 TT 倍输入复用

一个 T×TT\times T 输出块在某个 KK 分块上读取两个 T×TT\times T 输入块,却完成 T3T^3 次标量乘加。

每个载入的元素会被用于约 TT 个输出元素。

图 6

在方阵与理想复用的简化模型下,外部读流量由 O(N³) 降到约 O(N³/TILE)。

原视频 · 03:20 ↗

在方阵、整除、忽略边界与写回等理想假设下,输入 HBM 流量可概括为

O ⁣(N3T).O\!\left(\frac{N^3}{T}\right).

相对朴素 O(N3)O(N^3) 模型,约减少 TT 倍。

真实 kernel 的 tile 大小还受 shared memory 容量、寄存器压力、occupancy、bank conflict 和指令布局约束。

跟练与练习

原视频跟练

  • 先暂停在 00:40,手写 CijC_{ij} 的求和式,并标出收缩维 kk
  • 回到 01:40,在 4×44\times4 tile 网格中圈出一个 CC tile 对应的 AA tile 行和 BB tile 列。
  • 从 02:20 开始复述一次“HBM → shared memory → 寄存器 → HBM”的数据生命周期。
  • 在 03:20 的复杂度比较中逐项写出隐含假设,不把大 OO 当作实测带宽。
编者练习

A,BA,B 都是 64×6464\times64 方阵,tile 边长 T=16T=16。 只按视频的理想模型估算:相对完全不复用的输入读取,分块能带来约多少倍复用?输出矩阵共有多少个 tile?

查看参考答案

理想输入流量从 O(N3)O(N^3) 变为 O(N3/T)O(N^3/T),因此约减少 T=16T=16 倍。
每个维度的 tile 数为 64/16=464/16=4,输出共有 4×4=164\times4=16 个 tile。
该答案没有计入写回、边界、cache 和硬件资源约束,不能直接解释为实际 kernel 一定快 16 倍。

常见误区

  • 把 GEMM 固定写成 C=ABC=AB:完整接口还有 α\alphaβ\beta 与旧 CC
  • CC 的输入值和输出值混成一个对象:推导时可写 CinC_{\text{in}}CoutC_{\text{out}}
  • 把 FLOPs 等同于性能:相同 FLOPs 下,HBM 流量与数据复用会显著影响速度。
  • 把所有片上 SRAM 都称为 shared memory:shared memory 只是软件可管理的片上存储空间之一。
  • 认为 tile 越大越好:tile 过大会增加寄存器和 shared memory 压力,降低 occupancy。
  • O(N3/T)O(N^3/T) 当成精确字节数:它来自视频的理想化方阵访存模型。

本课小结

  • GEMM 的完整形式是 Cout=αAB+βCinC_{\text{out}}=\alpha AB+\beta C_{\text{in}}
  • CijC_{ij}AA 的第 ii 行与 BB 的第 jj 列的点积。
  • 分块的核心不是减少乘加次数,而是让输入 tile 在片上被反复使用。
  • shared memory 承担线程块级复用,寄存器保存细粒度输出累加器。
  • 视频理想模型中,边长为 TT 的 tile 把输入 HBM 流量从 O(N3)O(N^3) 降到 O(N3/T)O(N^3/T)
  • 真实 GEMM 还要在 tile 大小、占用率、同步与存储冲突之间权衡。
02

主题讲解 · 02:33

列并行与行并行究竟怎样切分矩阵

学习目标

  • 为线性层 Y=XWY=XW 建立不含歧义的 shape 账本。
  • 按数学矩阵 WW 的列或行判断 column parallel 与 row parallel。
  • 推导两种切分下每张卡的局部输入、权重与输出 shape。
  • 区分输出拼接与部分和归约两类集合通信。
  • 说明为什么实际系统可以延迟 all-gather,或改用 reduce-scatter。
  • 识别权重转置存储导致的命名错觉。

前置与衔接

上一课把 GEMM 写成矩阵 shape 与 tile 数据流。

本课把同一个乘法分给 pp 张设备。

统一采用

XRM×K,WRK×N,Y=XWRM×N.X\in\mathbb R^{M\times K},\qquad W\in\mathbb R^{K\times N},\qquad Y=XW\in\mathbb R^{M\times N}.

只要始终盯住收缩维 KK 与输出维 NN,列并行和行并行就不会混。

核心讲解

1. “列”与“行”首先指数学表达式中的 WW

tensor parallel 把同一层权重分片到多张卡上。

图 1

列并行与行并行都以权重 W 的切分方向命名,并分别产生输出分片或局部和。

原视频 · 00:00 ↗

这里的命名基于纸面公式 Y=XWY=XWWW 的方向:

  • 沿 WW 的列,也就是输出维 NN 切分,叫 column parallel;
  • 沿 WW 的行,也就是输入/收缩维 KK 切分,叫 row parallel。

某些框架实际保存的是 WW^\top,视觉上的行列可能反过来。

所以判断时应先把实现还原到数学 shape,而不是只看内存里矩形朝向。

2. 列并行切的是输出特征

假设 NN 能被设备数 pp 整除,把权重写成横向拼接:

W=[W0  W1    Wp1],WrRK×N/p.W=[W_0\;W_1\;\cdots\;W_{p-1}], \qquad W_r\in\mathbb R^{K\times N/p}.
图 2

列并行沿 W 的输出维切成 W1、W2,每张卡只保存一半列参数。

原视频 · 00:40 ↗

输入 XX 在各设备上保持完整或逻辑复制。

rr 张卡计算

Yr=XWrRM×N/p.Y_r=XW_r\in\mathbb R^{M\times N/p}.

每张卡得到最终输出的一组列,而不是同一个输出的部分和。

3. 列并行的数学合并是 concat

完整输出为

Y=[Y0  Y1    Yp1].Y=[Y_0\;Y_1\;\cdots\;Y_{p-1}].
图 3

列并行复制完整 X,各卡计算 XW1、XW2,结果是完整 Y 的左右分片,逻辑拼接即可复原。

原视频 · 01:00 ↗

若下一层必须在每张卡上看到完整 YY,可以用 all-gather 物理收集这些分片。

但“数学上可 concat”不等于“此刻一定发起 all-gather”。

若后续算子能直接消费分片,例如逐元素激活或配套的 row-parallel 层,就可以保留 sharded layout,延迟甚至省掉这次通信。

4. 行并行切的是收缩维

WW 沿行方向写成纵向堆叠:

W=[W0W1Wp1],WrRK/p×N.W= \begin{bmatrix} W_0\\W_1\\\vdots\\W_{p-1} \end{bmatrix}, \qquad W_r\in\mathbb R^{K/p\times N}.

此时输入也必须按相同的 KK 维切成

X=[X0  X1    Xp1],XrRM×K/p.X=[X_0\;X_1\;\cdots\;X_{p-1}], \qquad X_r\in\mathbb R^{M\times K/p}.
图 4

行并行沿 W 的输入维切分,因此 X 也必须沿对应特征维切成 X1、X2 才能满足内维对齐。

原视频 · 01:40 ↗

这样局部矩阵乘 XrWrX_rW_r 的收缩维才一致。

5. 行并行产生的是同形状部分和

每张卡计算

Pr=XrWrRM×N.P_r=X_rW_r\in\mathbb R^{M\times N}.
图 5

各卡得到 X1W1 与 X2W2,它们 shape 都等于完整 Y,但数值只是沿收缩维的局部和。

原视频 · 02:00 ↗

这些 PrP_r 不是不同输出列,而是完整 YY 在不同 KK 区间上的贡献。

因为

XW=r=0p1XrWr,XW = \sum_{r=0}^{p-1}X_rW_r,

所以必须逐元素相加,而不是沿某一维拼接。

6. All-Reduce 完成全局求和

若每张卡都需要完整输出,则执行

Y=AllReducesum(Pr).Y=\operatorname{AllReduce}_{\text{sum}}(P_r).
图 6

对局部和做逐元素 All-Reduce,得到 Y=X1W1+X2W2 的完整结果。

原视频 · 02:20 ↗

All-Reduce 同时完成 reduce 与结果分发,每个参与者最后都得到相同的 YY

若后续只想保留输出分片,工程实现也可以使用 reduce-scatter,把“求和”和“按目标布局分片”合并。

这属于编者补充的系统边界;视频用 All-Reduce 讲清最基本关系。

7. 用 shape 账本一眼区分两种模式

模式WrW_rXrX_rXX局部结果数学合并
column parallelK×N/pK\times N/pM×KM\times KM×N/pM\times N/p沿 NN concat
row parallelK/p×NK/p\times NM×K/pM\times K/pM×NM\times N对设备求和

一个实用判断法是:

  • 局部输出的列数缩小了,通常是列并行;
  • 局部输出 shape 完整但数值不完整,通常是行并行。

8. 通信需求取决于上下游布局

孤立分析一个线性层时,很容易说“列并行后 all-gather,行并行后 all-reduce”。

系统优化真正关心的是连续算子的布局:

  • 当前输出是不是下一层恰好需要的输入分片;
  • 逐元素算子能否就地作用于分片;
  • 结果要求复制到所有设备,还是继续保持切分;
  • forward 与 backward 各自需要什么 collective。

因此并行策略必须按算子链设计,而不是逐层机械插通信。

跟练与练习

原视频跟练

  • 从 00:40 开始,给每个 column shard 写出 K×N/pK\times N/p
  • 在 01:00 暂停,判断局部结果是最终输出的“分片”还是“部分和”。
  • 从 01:40 开始,检查 row parallel 为什么必须同步切分 XX 的特征维。
  • 在 02:20 解释 All-Reduce 为什么是求和而不是 concat。
编者练习

给定 XR8×12X\in\mathbb R^{8\times12}WR12×20W\in\mathbb R^{12\times20},使用 p=4p=4 张卡。 分别写出 column parallel 与 row parallel 下单卡权重、输入和局部输出的 shape,并指出合并操作。

查看参考答案

column parallel:
Wr:12×5,X:8×12,Yr:8×5.W_r:12\times5,\quad X:8\times12,\quad Y_r:8\times5.
四个 YrY_r 沿输出维 concat,得到 8×208\times20
row parallel:
Wr:3×20,Xr:8×3,Pr:8×20.W_r:3\times20,\quad X_r:8\times3,\quad P_r:8\times20.
四个 PrP_r 逐元素求和;若所有卡都要完整结果,可用 All-Reduce。

常见误区

  • 只看存储里的矩阵朝向命名行列:实现可能保存 WW^\top,先还原数学 shape。
  • 列并行也把 XX 沿 KK 切开:那会让局部结果变成部分和,已经进入行并行逻辑。
  • 把列并行的输出相加:各卡拥有不同输出列,数学上应 concat。
  • 把行并行的部分和 concat:每份都是 M×NM\times N 的同形状贡献,应逐元素相加。
  • 认为列并行后必定立即 all-gather:下游能消费分片时可以延迟通信。
  • 认为 row parallel 永远只能 All-Reduce:目标布局为分片时可考虑 reduce-scatter。

本课小结

  • Y=XWY=XW,column parallel 沿 WW 的输出维 NN 切分。
  • 列并行复制输入,产生 M×N/pM\times N/p 输出分片,数学合并是 concat。
  • row parallel 沿 WW 的收缩维 KK 切分,输入必须按同一维对齐切分。
  • 行并行产生 M×NM\times N 的局部部分和,数学合并是求和。
  • All-gather、All-Reduce 或 reduce-scatter 的选择取决于下游所需布局。
  • 权重转置存储会制造行列错觉,shape 账本比视觉方向可靠。
03

主题讲解 · 03:38

为什么 Ring All-Reduce 的单卡通信量趋近常数

学习目标

  • 区分 reduce、broadcast、reduce-scatter、all-gather 与 All-Reduce。
  • 解释中心化归约为何让根节点成为通信热点。
  • 追踪 ring reduce-scatter 中一个 chunk 的逐跳累加。
  • 解释 all-gather 如何把归约完成的 chunk 分发给所有设备。
  • 推导 NN 卡 ring All-Reduce 的单卡发送量。
  • 正确解读“4 卡 1.5 倍、128 卡约 2 倍”中的统计口径与延迟边界。

前置与衔接

上一课看到 row parallel 产生多个同形状部分和,必须跨设备求和。

All-Reduce 的语义是:

y=r=0N1xr,y=\sum_{r=0}^{N-1}x_r,

并让每个 rank 最终都拥有 yy

本课不只问“结果是什么”,还问这件事怎样避免让一张卡承担几乎全部流量。

核心讲解

1. 中心化方案把归约和广播都压在根节点

最直观的方法是选择一个 root:

  1. 其余 N1N-1 张卡把大小为 SS 的张量发给 root;
  2. root 对所有张量求和;
  3. root 把结果再发给其余 N1N-1 张卡。
图 1

朴素主从 All-Reduce 先把其他卡数据汇聚到主卡,再由主卡广播总和。

原视频 · 00:20 ↗

算法语义正确,但 root 同时接收和发送大量数据,网络链路和本地注入带宽都可能成为瓶颈。

2. 根节点负担会随设备数线性增长

按视频口径,root 在归约阶段接收 (N1)S(N-1)S,在广播阶段发送 (N1)S(N-1)S

总端点流量为

Vroot=2(N1)S.V_{\text{root}}=2(N-1)S.
图 2

按视频的传输量口径,主卡承担 reduce 与 broadcast 两阶段共 2(N−1)S 的通信。

原视频 · 00:40 ↗

非 root 设备各自只发送一次、接收一次,负载高度不均。

设备越多,root 的流量越大,这正是 ring 方案要消除的热点。

3. Ring 先把每张卡的数据切成 NN 个 chunk

NN 个 rank 连成逻辑环,每份大小为 SS 的输入切成 NN 个等大 chunk:

xr=[xr(0),xr(1),,xr(N1)],xr(j)=S/N.x_r=[x_r^{(0)},x_r^{(1)},\ldots,x_r^{(N-1)}], \qquad |x_r^{(j)}|=S/N.
图 3

Ring All-Reduce 把每卡大小 S 的张量等分成 N 个大小 S/N 的同编号数据块。

原视频 · 01:00 ↗

每一步,每个 rank 向下一个邻居发送一个 chunk,同时从上一个邻居接收一个 chunk。

所有 rank 都在并行通信,不再存在唯一 root。

4. Reduce-Scatter 让同编号 chunk 沿环累加

第一阶段需要 N1N-1 步。

收到一个 chunk 后,rank 把它与自己对应编号的局部 chunk 求和,再把部分和继续传向下一站。

图 4

Reduce-Scatter 中,同编号 chunk 沿环传递并逐卡累加,经过 N−1 步后每卡持有一个完整归约块。

原视频 · 01:40 ↗

阶段结束时:

  • 每个编号的 chunk 都已经聚合了全部 NN 张卡的贡献;
  • 每个 rank 只持有其中一个完成归约的 chunk;
  • 所有完成 chunk 分散在不同 rank 上。

这就是名字中的 reduce 与 scatter。

5. All-Gather 再分发已经归约完成的 chunk

第二阶段同样执行 N1N-1 步,但不再做求和,只转发完整 chunk。

图 5

All-Gather 再传播各卡已归约完成的 chunk,使每张卡最终收集全部 N 个结果块。

原视频 · 02:20 ↗

每一步,每个 rank 发送自己当前持有的一份完成 chunk,并接收另一份。

N1N-1 步之后,每个 rank 都收齐 NN 份归约 chunk,拼成完整结果 yy

因此 ring All-Reduce 可分解为:

Reduce-Scatter+All-Gather.\text{Reduce-Scatter}+\text{All-Gather}.

6. 单卡发送量趋近 2S2S

每个阶段有 N1N-1 步,每步每个 rank 发送大小 S/NS/N 的 chunk。

因此单个 rank 在两个阶段发送的总字节数为

Vsend=2(N1)SN=2(11N)S.V_{\text{send}} =2(N-1)\frac{S}{N} =2\left(1-\frac1N\right)S.
图 6

两阶段每卡发送量为 2(N−1)S/N:4 卡是 1.5S,128 卡是 1.984375S,趋近 2S。

原视频 · 03:00 ↗

代入 N=4N=4

Vsend=2×34S=1.5S.V_{\text{send}}=2\times\frac34S=1.5S.

代入 N=128N=128

Vsend=2×127128S=1.984375S.V_{\text{send}}=2\times\frac{127}{128}S=1.984375S.

NN\to\infty 时,单卡发送量趋近 2S2S,而不是随 NN 无界增长。

7. 必须先声明通信量统计口径

上式统计的是“每个 rank 发送的总字节数”。

同一个 rank 在环中也会接收同样多的字节。

若把端点的发送与接收相加,口径会变成

Vsend+recv=4(11N)S.V_{\text{send+recv}} =4\left(1-\frac1N\right)S.

所以“约 2 倍”与“约 4 倍”可能都出现,差别在于是否把收和发分开统计。

比较算法或 profiler 数据前,必须确认 SS 是单份张量大小、发送量、链路流量还是收发合计。

8. 字节数趋近常数,不代表扩容没有代价

ring 的两个阶段各有 N1N-1 个通信步,总步数为

2(N1).2(N-1).

单卡总字节数趋近常数,但启动延迟项会随步数增长。

若一次通信时间近似写成

T2(N1)α+2N1NSβ,T\approx 2(N-1)\alpha +2\frac{N-1}{N}S\beta,

α\alpha 表示每步固定延迟,β\beta 表示单位字节传输时间。

这是编者补充的常用性能模型,视频主要讲字节量。

跨节点链路、环顺序、网络拓扑、chunk 大小和并发也都会影响真实速度。

跟练与练习

原视频跟练

  • 从 00:20 开始列出中心化方案中 root 的接收与发送次数。
  • 在 01:00 暂停,把每卡张量标成 NN 个编号 chunk。
  • 从 01:40 追踪一个编号的 chunk 经过 N1N-1 步后聚合了哪些 rank。
  • 从 02:20 区分“继续累加”和“只做转发”的阶段边界。
  • 在 03:00 手算 4 卡与 128 卡的 2(N1)S/N2(N-1)S/N
编者练习

8 张卡对大小为 S=256 MiBS=256\text{ MiB} 的张量执行 ring All-Reduce。 按视频的单 rank 发送量口径,计算每张卡总共发送多少 MiB;若把收与发相加,又是多少 MiB?

查看参考答案

发送量为
2(118)×256=448 MiB.2\left(1-\frac18\right)\times256 =448\text{ MiB}.
接收量相同,所以端点收发合计为
448+448=896 MiB.448+448=896\text{ MiB}.
这里没有计入协议开销,也没有用步数延迟估算耗时。

常见误区

  • 把 All-Reduce 当成只求和到一个 root:最终每个参与 rank 都得到归约结果。
  • 认为 reduce-scatter 结束后每卡已有完整张量:每卡只有一个完整归约 chunk。
  • 认为 all-gather 还在做数值求和:第二阶段只分发已完成的 chunk。
  • 把单卡发送量与收发合计混用:两种口径相差一倍。
  • 认为设备越多单卡字节数越少到零:它趋近 2S2S
  • 由字节量趋近常数推出延迟不变:通信步数仍为 2(N1)2(N-1)

本课小结

  • 中心化 reduce+broadcast 让 root 承担 2(N1)S2(N-1)S 的热点流量。
  • ring 把张量切成 NN 个 chunk,并让所有 rank 同时参与邻居通信。
  • reduce-scatter 完成分块归约,all-gather 完成结果分发。
  • 单 rank 发送量为 2(11/N)S2(1-1/N)S:4 卡是 1.5S1.5S,128 卡是 1.984375S1.984375S
  • 收发合计需再乘二,比较数字前必须统一统计口径。
  • 单卡字节数趋近常数不等于延迟恒定,步数与拓扑仍很重要。
04

主题讲解 · 03:52

用 Ping-Pong 双缓冲隐藏 GEMM 搬运延迟

学习目标

  • 识别 tiled GEMM 中沿 KK 维依次消费的输入 tile 对。
  • 说明单缓冲“先搬后算”为什么让计算单元等待。
  • 按 prologue、steady state、epilogue 描述双缓冲流水线。
  • 正确区分 ping buffer、pong buffer 与寄存器输出累加器。
  • 写出搬运与计算重叠后的理想阶段耗时。
  • 指出异步拷贝、同步与资源占用对双缓冲收益的边界。

前置与衔接

P46 已经说明:一个输出 tile 要沿收缩维 KK 依次加载多对 AABB tile。

若目标矩阵为 128×128128\times128、tile 边长为 32,一个输出 tile 要经历四个 KK 分块。

本课进一步追问:每轮都要从 HBM 搬下一对 tile,能否让搬运与当前计算同时发生?

答案是 ping-pong 双缓冲,但“有两个 buffer”只是起点,真正关键是正确流水。

核心讲解

1. 一个输出 tile 沿 KK 维消费多对输入 tile

目标块 CI,JC_{I,J} 的计算可写为

CI,J=k=0K/T1AI,kBk,J.C_{I,J} =\sum_{k=0}^{K/T-1}A_{I,k}B_{k,J}.
图 1

一个 C 目标 tile 沿 K 维依次消费四对 A/B tile,并把乘积累加到同一结果块。

原视频 · 00:20 ↗

每一轮都要:

  1. 从 HBM 读取一对 AABB tile;
  2. 放入 shared memory;
  3. 线程执行局部乘加;
  4. 把结果累加到寄存器中的 CC fragment。

输出累加器跨越所有 KK 分块持续存在。

2. 单缓冲让 Load 与 Compute 串行

若只有一块 shared memory 区域,当前计算读取它时,下一批数据不能安全覆盖同一区域。

最直接的时间线是

L0,C0,L1,C1,L2,C2,L_0,C_0,L_1,C_1,L_2,C_2,\ldots

其中 LiL_i 是第 ii 批加载,CiC_i 是对应计算。

图 2

单缓冲串行执行“加载→计算→再加载”,数据搬运期间计算单元会等待。

原视频 · 01:20 ↗

每一轮耗时近似为

Tstage=TL+TC.T_{\text{stage}}=T_L+T_C.

加载时计算单元空闲,计算时搬运通道也可能没有为下一轮工作。

3. Ping 与 Pong 提供两个可交替占用的槽位

双缓冲在 shared memory 中准备两套区域:ping 和 pong。

第一步 prologue 只能先把第 0 批数据加载到 ping,因为此时还没有可计算的数据。

图 3

SRAM 被分为 ping/pong 两区;流水线序幕 T0 先把首对 tile 载入 ping。

原视频 · 01:40 ↗

prologue 的时间线是

L0ping.L_0\to\text{ping}.

视频 ASR 出现了近似“pink”“pom”的音,画面和上下文对应的标准术语是 ping/pong。

4. 稳态阶段让当前计算与下一批加载重叠

ping 中已有第 0 批数据后,可以同时:

  • 从 ping 读取并执行 C0C_0
  • 把第 1 批数据异步加载到 pong。
图 4

稳态中一块缓冲区供计算,另一块并行预取下一对 tile,随后交换角色。

原视频 · 02:20 ↗

如果两条路径确实能并发,稳态阶段耗时理想化为

Tstagemax(TL,TC),T_{\text{stage}}\approx\max(T_L,T_C),

而不是 TL+TCT_L+T_C

所谓“隐藏搬运延迟”表示较短的那部分被较长的那部分覆盖,不表示搬运时间物理消失。

5. 两个槽位逐轮交换读写角色

第 0 轮完成后,第 1 批已在 pong 中。

下一阶段改为:计算 pong 中的第 1 批,同时把第 2 批加载到 ping。

图 5

T2/T3 阶段 ping 与 pong 交替承担计算和加载,寄存器中的 C 累加器持续保留。

原视频 · 02:40 ↗

稳态时间线可以写成:

阶段计算读取异步加载
1ping 中第 0 批第 1 批到 pong
2pong 中第 1 批第 2 批到 ping
3ping 中第 2 批第 3 批到 pong

切换前必须确认下一批加载完成,复用槽位前也必须确认上一轮读取完成。

6. Epilogue 只剩最后一次计算与输出写回

最后一批数据已经加载后,不再有下一批可预取。

流水线进入 epilogue:完成最后一次计算,然后把寄存器中完整的 CC tile 写回 HBM。

图 6

最后一对 tile 完成后,寄存器中的目标 C tile 才一次性写回 HBM。

原视频 · 03:20 ↗

完整结构因此是:

prologue loadoverlapped steady statesepilogue compute/writeback.\text{prologue load} \rightarrow \text{overlapped steady states} \rightarrow \text{epilogue compute/writeback}.

首批加载和末批计算无法被相邻工作完全隐藏,这称为流水线填充与排空开销。

7. 双缓冲正确运行需要三个条件

第一,硬件和指令路径支持传输与计算并发。

例如需要真正的异步拷贝或独立 copy pipeline;把两个同步操作顺序写在代码中不会自动重叠。

第二,必须有正确同步。

消费者不能在 load 完成前读取 buffer,生产者也不能在旧数据尚被使用时覆盖它。

第三,计算时间要足以覆盖搬运时间。

TL>TCT_L>T_C,理想稳态仍需 TLT_L;双缓冲只能隐藏 TCT_C 对应的那一部分等待。

8. 空间换时间也会改变 occupancy

双缓冲让 shared memory 占用从一套 tile 增加到两套。

额外资源可能减少一个 SM 同时驻留的线程块数量。

收益应比较:

  • 延迟重叠节省了多少时间;
  • shared memory 与寄存器压力增加多少;
  • occupancy 下降是否削弱吞吐;
  • tile 边界和同步开销是否显著。

现代高性能 GEMM 还可能使用多级流水,而不只两个 stage;这是编者补充的实现边界。

跟练与练习

原视频跟练

  • 从 00:20 开始给四对输入 tile 标记 0,1,2,30,1,2,3
  • 在 01:20 暂停,把单缓冲时间线写成 L0,C0,L1,C1L_0,C_0,L_1,C_1
  • 从 01:40 识别 prologue 为什么只有加载没有计算。
  • 在 02:20 说明当前 buffer 与下一 buffer 为什么不能是同一物理槽位。
  • 从 02:40 画出 ping/pong 连续三阶段的角色交换。
  • 在 03:20 指出 epilogue 中哪一项不能再被下一批加载覆盖。
编者练习

某 tile 的 HBM 加载用时 TL=8μsT_L=8\,\mu s,计算用时 TC=12μsT_C=12\,\mu s,共有 4 个 KK 分块。 忽略同步与写回,比较串行方案与理想双缓冲方案的总时间。

查看参考答案

串行方案为
4(TL+TC)=4×20=80μs.4(T_L+T_C)=4\times20=80\,\mu s.
理想双缓冲包含一次 prologue load、三个重叠间隔与最后一次计算,可写为
TL+(41)max(TL,TC)+TC=8+3×12+12=56μs.T_L+(4-1)\max(T_L,T_C)+T_C =8+3\times12+12 =56\,\mu s.
这是理想上界估算;真实结果还受同步、资源与异步拷贝能力影响。

常见误区

  • 认为准备两个数组就会自动并发:必须有硬件支持的异步传输与正确调度。
  • 把 ping/pong 当成两个输出累加器:它们缓存输入 tile,输出部分和通常在寄存器。
  • 让 load 覆盖正在计算的槽位:角色交换需要完成事件或屏障保护。
  • 认为搬运时间被完全消灭:稳态仍由 max(TL,TC)\max(T_L,T_C) 决定。
  • 忽略 prologue 与 epilogue:第一批和最后一批存在不可重叠边界。
  • 认为双缓冲没有成本:它会增加 shared memory 占用并可能降低 occupancy。

本课小结

  • tiled GEMM 沿 KK 维连续消费多对输入 tile,并在寄存器中累加输出。
  • 单缓冲把 load 与 compute 串行,阶段耗时为 TL+TCT_L+T_C
  • ping-pong 双缓冲让当前计算与下一批加载使用不同槽位。
  • 理想稳态耗时下降为 max(TL,TC)\max(T_L,T_C),但仍有填充和排空开销。
  • 正确重叠依赖异步 copy、同步和足够的计算覆盖窗口。
  • 双缓冲以更多片上空间换取时间,必须结合 occupancy 评估整体收益。
05

主题讲解 · 03:33

Megatron 如何用先列后行压缩 FFN 通信

学习目标

  • 复述 column-parallel 与 row-parallel 线性层的 shape 规则。
  • 为两层 FFN 写出 D4DDD\to4D\to D 的中间张量 shape。
  • 推导两张卡上 up projection 的列分片。
  • 解释逐元素 GELU 为什么能直接作用于分片而无需 all-gather。
  • 推导 down projection 的行分片与最终部分和 All-Reduce。
  • 说明“只通信一次”在前向、网络结构和现代实现中的适用边界。

前置与衔接

P47 分别讨论了列并行和行并行。

若孤立看每个线性层,可能会在列并行后立刻 all-gather,又在行并行后 All-Reduce。

Megatron 风格 tensor parallel 的关键是成对安排两层:

column-parallel uplocal activationrow-parallel down.\text{column-parallel up} \rightarrow \text{local activation} \rightarrow \text{row-parallel down}.

中间分片直接衔接,从而避免一次不必要的聚合。

核心讲解

1. 先回顾 column-parallel 线性层

Y=XW,XRM×K,WRK×N,Y=XW, \quad X\in\mathbb R^{M\times K}, \quad W\in\mathbb R^{K\times N},

column parallel 沿输出维 NN 切分 WW

W=[W0  W1    Wp1].W=[W_0\;W_1\;\cdots\;W_{p-1}].
图 1

列并行切 W 的输出维,复制 X,各卡得到可沿特征维拼接的输出分片。

原视频 · 00:20 ↗

输入 XX 复制到各卡,局部结果为

Yr=XWrRM×N/p.Y_r=XW_r\in\mathbb R^{M\times N/p}.

YrY_r 是完整输出的不同特征分片。

2. 再回顾 row-parallel 线性层

row parallel 沿输入/收缩维 KK 切分权重,并要求输入也按同一维切分:

W=[W0W1Wp1],X=[X0  X1    Xp1].W= \begin{bmatrix}W_0\\W_1\\\vdots\\W_{p-1}\end{bmatrix}, \qquad X=[X_0\;X_1\;\cdots\;X_{p-1}].
图 2

行并行切 W 的输入维和 X 的对应特征分片,各卡生成同形局部和并通过归约相加。

原视频 · 01:00 ↗

每张卡产生同形状部分和

Pr=XrWrRM×N,P_r=X_rW_r\in\mathbb R^{M\times N},

完整输出是 Y=rPrY=\sum_rP_r

3. 标准两层 FFN 先扩张再压回隐藏维

视频采用简化的 Transformer FFN:

H=GELU(OWup),H=\operatorname{GELU}(OW_{\text{up}}),
Y=HWdown.Y=HW_{\text{down}}.

shape 为

ORM×D,WupRD×4D,WdownR4D×D.O\in\mathbb R^{M\times D}, \quad W_{\text{up}}\in\mathbb R^{D\times4D}, \quad W_{\text{down}}\in\mathbb R^{4D\times D}.

MM 可代表 batch、sequence 等非隐藏维展平后的行数。

中间激活 HH 的 shape 是 M×4DM\times4D

4. Up projection 用列并行产生中间特征分片

以两张卡为例,把 WupW_{\text{up}}4D4D 个输出列均分:

Wup=[Wup,0  Wup,1],W_{\text{up}}=[W_{\text{up},0}\;W_{\text{up},1}],
Wup,rRD×2D.W_{\text{up},r}\in\mathbb R^{D\times2D}.
图 3

W_up 从 D 扩到 4D,按列切到两卡后,每卡只产生宽 2D 的中间激活分片。

原视频 · 02:00 ↗

两张卡都接收完整 ORM×DO\in\mathbb R^{M\times D},分别得到

Zr=OWup,rRM×2D.Z_r=OW_{\text{up},r}\in\mathbb R^{M\times2D}.

逻辑上,完整 ZZ[Z0  Z1][Z_0\;Z_1],但此处不急着 all-gather。

5. GELU 是逐元素算子,可以在各分片独立执行

GELU 对每个标量独立作用:

Hr=GELU(Zr).H_r=\operatorname{GELU}(Z_r).
图 4

GELU 是逐元素操作,可直接在每卡的 2D 分片上独立执行,无需先 All-Gather 成 4D。

原视频 · 02:20 ↗

因为 HH 的一个特征元素不依赖其他特征分片,满足

GELU([Z0  Z1])=[GELU(Z0)  GELU(Z1)].\operatorname{GELU}([Z_0\;Z_1]) =[\operatorname{GELU}(Z_0)\;\operatorname{GELU}(Z_1)].

所以无需先拼成完整 M×4DM\times4D 再做激活。

这一步省掉了 column-parallel 层后原本可能出现的 all-gather。

若中间算子跨隐藏维归约或混合特征,该结论就不再自动成立。

6. Down projection 用行并行直接接住激活分片

WdownR4D×DW_{\text{down}}\in\mathbb R^{4D\times D} 沿行切成两份:

Wdown=[Wdown,0Wdown,1],W_{\text{down}}= \begin{bmatrix} W_{\text{down},0}\\ W_{\text{down},1} \end{bmatrix},
Wdown,rR2D×D.W_{\text{down},r}\in\mathbb R^{2D\times D}.
图 5

W_down 按输入维行切,每卡用本地 2D 激活乘本地 2D×D 权重,得到宽 D 的局部和。

原视频 · 03:00 ↗

它与本卡已有的 HrRM×2DH_r\in\mathbb R^{M\times2D} 正好对齐:

Pr=HrWdown,rRM×D.P_r=H_rW_{\text{down},r} \in\mathbb R^{M\times D}.

这就是“先列后行”能无缝连接的 shape 原因。

7. 最终只需把 down projection 的部分和相加

两个局部结果 P0,P1P_0,P_1 都是完整输出 shape,但各自只包含一半中间特征的贡献。

因此

Y=P0+P1.Y=P_0+P_1.

若每张卡都需要完整 YY,执行求和 All-Reduce。

图 6

两卡局部 D 维结果在 FFN 末端做一次 All-Reduce,恢复完整输出。

原视频 · 03:20 ↗

前向主干于是只有 FFN 末端这一次必要同步,而中间没有 all-gather。

“减少通信”的本质不是改变线性层数学结果,而是让第一层的输出布局成为第二层恰好需要的输入布局。

8. “只通信一次”有明确边界

首先,这里说的是视频所示简化 FFN 的前向路径。

反向传播仍需与梯度布局相匹配的 collective,不能由前向图推出训练全程只有一次通信。

其次,现代 LLM 常使用 gated FFN,例如 SwiGLU:

H=SiLU(OWg)(OWu).H=\operatorname{SiLU}(OW_g)\odot(OW_u).

通常会把 gate projection 与 up projection 按相同输出维分组切分,使逐元素乘仍能局部执行,再接 row-parallel down projection。

这是编者对现代结构的补充;视频讲的是单个 GELU 分支。

再次,sequence parallel 等布局可能把末端 All-Reduce 替换或拆成 reduce-scatter 与后续 all-gather。

通信原语会变化,但“让相邻算子的分片布局对齐”仍是核心原则。

9. Bias 的放置也要服从部分和语义

若 down projection 带 bias bRDb\in\mathbb R^D,不能让每张卡都在局部部分和上加完整 bb 后再求和。

否则 pp 张卡会得到

r(Pr+b)=rPr+pb.\sum_r(P_r+b)=\sum_rP_r+pb.

常见做法是归约后只加一次 bias,或采用能保证全局只贡献一次的分片约定。

这一点是编者补充,用来检验是否真正理解“局部结果是部分和”。

跟练与练习

原视频跟练

  • 从 00:20 回顾 column parallel 的输出分片 shape。
  • 从 01:00 回顾 row parallel 为什么产生同形状部分和。
  • 在 02:00 暂停,为两张卡写出 Wup,r:D×2DW_{\text{up},r}:D\times2D
  • 从 02:20 说明逐元素 GELU 为什么不要求跨卡数据。
  • 在 03:00 检查 Hr:M×2DH_r:M\times2DWdown,r:2D×DW_{\text{down},r}:2D\times D 的收缩维。
  • 从 03:20 解释最后一步为什么求和而不是 concat。
编者练习

M=16M=16D=1024D=1024,使用 4 张卡并行一个 D4DDD\to4D\to D 的 GELU FFN。 写出每张卡的 Wup,rW_{\text{up},r}ZrZ_rWdown,rW_{\text{down},r}PrP_r shape,并说明中间与末端是否需要通信。

查看参考答案

上投影采用列并行:
Wup,r:1024×1024,Zr:16×1024.W_{\text{up},r}:1024\times1024, \qquad Z_r:16\times1024.
GELU 逐元素执行,HrH_r 仍为 16×102416\times1024,中间不需要 all-gather。
下投影采用行并行:
Wdown,r:1024×1024,Pr:16×1024.W_{\text{down},r}:1024\times1024, \qquad P_r:16\times1024.
四个 PrP_r 是部分和;若每卡都要完整输出,末端做一次 sum All-Reduce。

常见误区

  • 列并行后机械执行 all-gather:GELU 和下一层的行分片都能直接消费当前分片。
  • ZrZ_r 当成完整中间激活:它只有 4D/p4D/p 个特征。
  • 让 down projection 继续按列切:这样无法直接接住现有的 HrH_r 输入分片。
  • 把末端部分和 concat:每个 PrP_r 都是 M×DM\times D,应逐元素求和。
  • 声称整个训练只有一次通信:视频结论针对简化 FFN 的这段前向路径。
  • 把 GELU 结论推广到任意中间算子:跨特征归约或混合会产生通信需求。
  • 在每张卡的部分和上都加完整 bias:All-Reduce 后会把 bias 重复 pp 次。

本课小结

  • Megatron 风格 FFN 把 up projection 设为 column parallel,把 down projection 设为 row parallel。
  • 上投影产生 M×4D/pM\times4D/p 中间特征分片。
  • GELU 逐元素局部执行,不改变分片布局,也无需 all-gather。
  • 下投影的行分片与本卡激活分片在收缩维上精确对齐。
  • 每卡得到 M×DM\times D 部分和,最终用 sum All-Reduce 合并。
  • “先列后行”的核心是布局连续性;反向、gated FFN、sequence parallel 和 bias 仍需单独审视。
06

主题讲解 · 02:48

外积为什么适合构造 GEMM 微内核

学习目标

  • 区分内积、外积以及它们对同一个 GEMM 的两种组织方式。
  • 从一个线程只算一个元素的基线说明输入复用不足。
  • 解释线程负责 4×44\times4 输出微块后为什么会增加寄存器需求。
  • 推导 4×14\times11×41\times4 外积怎样一次更新 16 个累加器。
  • 复现视频中 48 与 24 个标量槽位的概念计数。
  • 正确解读“四倍”是读数/乘加比改善,而非无条件的实测加速比。

前置与衔接

P46 已把 128×128128\times128 GEMM 切成边长为 32 的 tile,并让目标 CC tile 沿 KK 维累加四对 AABB tile。

P49 又用双缓冲把下一对 tile 的搬运和当前 tile 的计算重叠。

本课继续向微观层级放大:数据进入片上 SRAM 后,一个线程怎样用较少的输入寄存器持续更新一组输出?

核心变换不是改变公式

C=AB,C=AB,

而是把标量乘加按外积顺序组织。

核心讲解

1. 块级 GEMM 先把输入 tile 流进 SRAM

视频仍采用 128×128128\times128 方阵和 T=32T=32 的示例。

每个矩阵被切成 4×44\times4 个 tile。

为了完成一个目标 CC tile,需要沿 KK 方向依次载入四对 AABB tile。

图 1

128×128 GEMM 按 TILE=32 分块,A、B 的四对 tile 依次从 HBM 进入片上 SRAM 并累加目标 C tile。

原视频 · 00:20 ↗

SRAM 容量有限,所以旧 tile 在消费完后被新 tile 覆盖。

这解决了 HBM 层的数据复用,但 SRAM 到寄存器的取数方式仍会决定微内核效率。

2. 最朴素的内积让一个线程只负责一个输出元素

Cij=k=031AikBkj,C_{ij}=\sum_{k=0}^{31}A_{ik}B_{kj},

一个线程可以读取 AA 的一行和 BB 的一列,更新一个 CijC_{ij}

图 2

朴素标量思路让一个线程读取 A 的一行与 B 的一列,只更新 C 中一个元素。

原视频 · 00:40 ↗

这种分工容易理解,却让每个输入标量只服务于当前线程的一个输出。

相邻线程可能再次读取相同的 AA 行片段或 BB 列片段。

3. 标量基线的读数/乘加比很低

沿 K=32K=32 循环时,每轮读取一个 AikA_{ik} 与一个 BkjB_{kj},执行一次 fused multiply-add:

cc+a×b.c\leftarrow c+a\times b.
图 3

沿 K=32 循环时,标量内积每轮读两个数做一次乘加,32 轮读取 64 个标量并完成 32 次乘加。

原视频 · 01:00 ↗

按视频的标量计数:

Nread=32×2=64,N_{\text{read}}=32\times2=64,
NFMA=32.N_{\text{FMA}}=32.

因此每读取一个输入标量只得到

NFMANread=12\frac{N_{\text{FMA}}}{N_{\text{read}}}=\frac12

次乘加。

这里统计的是从 SRAM 到线程寄存器的逻辑标量读取,不是 profiler 中的 HBM 字节数。

4. 让线程负责 4×44\times4 输出可提高复用

若一个 32×3232\times32 tile 由 8×8=648\times8=64 个线程共同负责,每个线程可以拥有一个 4×44\times4 输出微块。

线程的累加器写成

C(t)R4×4.C^{(t)}\in\mathbb R^{4\times4}.

一共需要 16 个标量累加槽位。

同一份 AA 输入可横向复用到四个输出列,同一份 BB 输入可纵向复用到四个输出行。

5. 直接暂存两块 4×44\times4 输入会造成寄存器压力

视频先给出一种直接思路:同时把 AA4×44\times4 片段和 BB4×44\times4 片段放进寄存器,再更新 4×44\times4 输出。

图 4

若一个线程用直接内积同时更新 4×4 输出微块,需要保留 16 个累加器和 32 个输入标量,概念计数为 48。

原视频 · 01:40 ↗

按“一个标量槽位对应一个寄存器”的教学计数:

16  (C)+16  (A)+16  (B)=48.16\;(C)+16\;(A)+16\;(B)=48.

输入片段在完成当前更新后就不再需要,长期把 32 个输入标量同时留在寄存器中并不划算。

寄存器占用过高还可能降低一个 SM 同时驻留的 warp 数,甚至引发 spill。

6. 外积每轮只保留两条向量

KK 维循环拆成标量位置 kk

线程每轮读取

a=AI,kR4×1,a=A_{I,k}\in\mathbb R^{4\times1},
b=Bk,JR1×4.b=B_{k,J}\in\mathbb R^{1\times4}.

二者外积为

abR4×4,ab\in\mathbb R^{4\times4},

恰好覆盖线程负责的全部 16 个输出。

图 5

每轮读取 A 的 4×1 向量与 B 的 1×4 向量,外积一次更新完整 4×4 输出微块。

原视频 · 02:00 ↗

更新式可写为

Cuv(t)mathrel+=aubv,u,v{0,1,2,3}.C^{(t)}_{uv}mathrel{+}=a_u b_v, \qquad u,v\in\{0,1,2,3\}.

同一个 aua_u 被四个列位置复用,同一个 bvb_v 被四个行位置复用。

7. 外积把输入槽位从 32 个降到 8 个

每轮只需暂存 4 个 AA 标量和 4 个 BB 标量。

因此视频的概念寄存器计数为

16  (C)+4  (A)+4  (B)=24.16\;(C)+4\;(A)+4\;(B)=24.

相对直接暂存两个 4×44\times4 输入块的 48,标量槽位减半。

输入向量完成一次外积后即可被下一轮覆盖,只有 16 个输出累加器跨越全部 KK 循环长期存活。

8. 读数/乘加比提高四倍

每轮外积读取 4+4=84+4=8 个标量,完成 4×4=164\times4=16 次乘加。

循环 32 轮:

Nread=32×8=256,N_{\text{read}}=32\times8=256,
NFMA=32×16=512.N_{\text{FMA}}=32\times16=512.
图 6

外积方案每轮读取 8 个标量完成 16 次乘加,循环 32 轮共读取 256 个标量、完成 512 次乘加。

原视频 · 02:20 ↗

于是

NFMANread=2.\frac{N_{\text{FMA}}}{N_{\text{read}}}=2.

与标量基线的 1/21/2 相比,逻辑读数对应的乘加密度提高

21/2=4\frac{2}{1/2}=4

倍。

这就是视频所说“四倍”的准确口径。

9. 四倍算术密度不等于 kernel 必然快四倍

真实吞吐还受以下因素影响:

  • shared memory bank conflict 与访存合并;
  • 寄存器分配粒度和编译器向量化;
  • 指令吞吐、依赖链与调度;
  • occupancy、tile 边界和同步;
  • Tensor Core 或 CUDA Core 的具体数据通路。

视频中的 48 与 24 也是标量槽位教学模型。

物理寄存器分配由 ISA、数据类型、编译器和架构决定,不应把图中的数字当成所有实现的固定寄存器数。

跟练与练习

原视频跟练

  • 从 00:40 写出单个 CijC_{ij} 的行列内积。
  • 在 01:00 复算 32 轮为何读取 64 个标量、完成 32 次乘加。
  • 从 01:40 列出 16 个输出累加器与 32 个输入槽位。
  • 在 02:00 展开 4×14\times11×41\times4 外积的 16 个乘积。
  • 从 02:20 对比两种方案的 NFMA/NreadN_{\text{FMA}}/N_{\text{read}}
编者练习

一个线程负责 m×nm\times n 输出微块,每个 KK 位置使用一个 m×1m\times1 向量和一个 1×n1\times n 向量做外积。 忽略输出初始化,写出每轮输入标量数、乘加数,以及概念寄存器槽位数。

查看参考答案

每轮输入标量数为
m+n.m+n.
外积更新全部输出,乘加数为
mn.mn.
若输出累加器跨循环驻留、当前输入向量只保留一份,概念槽位数为
mn+m+n.mn+m+n.
例如 m=n=4m=n=4 时为 16+4+4=2416+4+4=24

常见误区

  • 把 ASR 的“外机”当成新硬件:这里指矩阵外积 outer product。
  • 把一个线程只算一个点当成唯一实现:高性能微内核常让线程持有多个输出累加器。
  • 认为外积减少了 GEMM 的 FLOPs:乘加总数不变,变化的是数据复用与执行顺序。
  • 把 48/24 当成精确物理寄存器数:它们是视频的标量槽位概念计数。
  • 把四倍直接解释为墙钟时间快四倍:四倍对应逻辑读数/乘加比。
  • 忽略 occupancy:微块过大虽提高复用,也会让输出累加器数量迅速增长。

本课小结

  • 内积按一个输出聚合 KK 维,外积按一个 KK 位置同时更新多个输出。
  • 4×14\times11×41\times4 外积一次产生 4×44\times4 的 16 个乘积。
  • 外积微内核只需 8 个当前输入槽位和 16 个长期累加器,视频计数为 24。
  • 标量基线每读一个输入完成 1/21/2 次乘加,外积方案提高到 2 次。
  • 视频中的四倍是逻辑算术密度比,不是无条件的真实吞吐加速。
  • 微块尺寸最终要在复用、寄存器压力、occupancy 与硬件指令之间折中。
07

主题讲解 · 03:49

两级双缓冲怎样串起 GEMM 数据流水线

学习目标

  • 区分 HBM→SRAM 与 SRAM→寄存器两条搬运路径。
  • 描述 SRAM 级 ping-pong 如何服务连续的 32×3232\times32 tile 对。
  • 32×3232\times32 输出 tile 推导线程负责的 8×88\times8 微块。
  • 解释寄存器级 ping-pong 如何预取下一轮 8×18\times11×81\times8 向量。
  • 复现 64 个累加槽位加 32 个输入槽位等于 96 的概念计数。
  • 指出双缓冲依赖异步搬运、同步和资源预算,不会自动隐藏全部延迟。

前置与衔接

P49 介绍了一层 ping-pong:当前 tile 计算时,把下一 tile 搬到另一个 SRAM 槽位。

P51 又说明外积微内核可以用两条向量更新一组输出累加器。

本课把二者叠起来:

HBMSRAM ping/pongregister ping/pongouter-product compute.\text{HBM} \rightarrow \text{SRAM ping/pong} \rightarrow \text{register ping/pong} \rightarrow \text{outer-product compute}.

每一层缓冲隐藏自己上游搬运的等待,但每层的粒度和生命周期并不相同。

核心讲解

1. 块级任务沿 KK 维消费四对大 tile

视频仍采用 128×128128\times128 方阵和 T=32T=32

一个红色目标 CC tile 的 shape 是 32×3232\times32

它需要 A 的一行四个 tile 与 B 的一列四个 tile:

CI,Jmathrel+=AI,kBk,J,k=0,1,2,3.C_{I,J}mathrel{+}=A_{I,k}B_{k,J}, \qquad k=0,1,2,3.
图 1

一个 32×32 目标 C tile 需要沿 K 维依次消费 A 的四个行 tile 与 B 的四个列 tile。

原视频 · 00:20 ↗

这一层迭代的单位是“一对 32×3232\times32 输入大块”。

2. SRAM 级双缓冲拥有两个大块槽位

把片上 SRAM 分成 ping 与 pong。

每个槽位能容纳一对 32×3232\times32AABB tile。

图 2

SRAM 被划成 ping 与 pong 两个槽位,每个槽位容纳一对 32×32 的 A、B tile。

原视频 · 00:40 ↗

prologue 先把第 0 对 tile 载入 ping。

此时 pong 为空,因为没有更早的计算可以覆盖第一次加载。

3. 当前计算与下一对 HBM 加载重叠

进入稳态后,可以同时执行:

  • 从 ping 读取第 0 对 tile 并计算;
  • 把第 1 对 tile 从 HBM 加载到 pong。
图 3

计算单元读取 ping 中当前 tile 对时,下一对橙色 tile 同时加载到 pong,随后两槽位换位。

原视频 · 01:00 ↗

下一阶段两个槽位交换角色:计算 pong,加载新数据到 ping。

理想稳态耗时近似

TSRAM-stagemax(THBM→SRAM,Ttile-compute).T_{\text{SRAM-stage}} \approx \max(T_{\text{HBM→SRAM}},T_{\text{tile-compute}}).

四对大 tile 全部消费后,寄存器中的完整 CC tile 才写回 HBM。

4. 只有 SRAM 双缓冲仍可能卡在寄存器取数

计算单元不能直接对整个 32×3232\times32 tile 一口气完成全部工作。

线程需要从 SRAM 继续读取更小片段到寄存器。

如果每轮都“先从 SRAM 读取,再做外积”,SRAM→寄存器延迟仍与计算串行。

因此第二级双缓冲位于线程寄存器输入槽位,而不是再复制一份完整输出累加器。

5. 一个线程负责 8×88\times8 输出微块

视频示例让一个线程负责目标 32×3232\times32 tile 中的一个 8×88\times8 区域。

它对应 A 的 8×328\times32 行片段与 B 的 32×832\times8 列片段。

图 4

在线程级放大图中,一个线程负责 8×8 的 C 微块,并使用 A 的 8×32 与 B 的 32×8 切片。

原视频 · 02:00 ↗

沿 K=32K=32 的每个位置,线程读取

akR8×1,bkR1×8,a_k\in\mathbb R^{8\times1}, \qquad b_k\in\mathbb R^{1\times8},

并执行

C(t)+=akbk.C^{(t)}\mathrel{+}=a_kb_k.

每轮外积产生 8×8=648\times8=64 次乘加。

6. 寄存器级 ping-pong 预取下一轮向量

寄存器输入槽位也分成 ping 与 pong。

第一轮从 SRAM 读取 a0,b0a_0,b_0 到 register ping。

稳态阶段同时执行:

  • Tensor/CUDA 计算路径用 ping 中的 ak,bka_k,b_k 更新 64 个 C 累加器;
  • load 路径把 ak+1,bk+1a_{k+1},b_{k+1} 预取到 register pong。
图 5

寄存器级 ping 正在提供 8×1 与 1×8 向量做外积时,下一轮向量被预取到 pong。

原视频 · 03:00 ↗

下一轮再交换 ping/pong。

理想稳态耗时近似

Treg-stagemax(TSRAM→reg,Touter-product).T_{\text{reg-stage}} \approx \max(T_{\text{SRAM→reg}},T_{\text{outer-product}}).

7. 两级流水的粒度不同

层级ping/pong 中的对象一次切换对应
SRAM一对 32×3232\times32 的 A、B tile一个大 KK tile
寄存器8×18\times1 的 A 向量与 1×81\times8 的 B 向量一个标量 kk 位置

SRAM 层解决 HBM 大块搬运,寄存器层解决片上小片段搬运。

两级缓冲嵌套,而不是两个彼此无关的技巧。

8. 视频的寄存器预算是 96 个标量槽位

8×88\times8 输出需要

8×8=648\times8=64

个累加槽位。

每个输入缓冲槽位需要 88 个 A 标量和 88 个 B 标量,共 16。

ping/pong 两套输入为

2×(8+8)=32.2\times(8+8)=32.

所以总计

64+32=96.64+32=96.
图 6

视频的概念计数包含 64 个 C 累加槽位和两套 A/B 输入槽位共 32 个,总计 96。

原视频 · 03:20 ↗

这比为两轮都复制 8×88\times8 输出累加器节省得多,因为 C 累加器不需要 ping-pong;它们始终是同一组长期状态。

9. 96 不是所有 GPU 上的精确物理寄存器数

视频把一个标量槽位对应为一个寄存器,适合解释空间关系。

真实分配还取决于:

  • 数据类型和寄存器宽度;
  • 编译器的 live range、向量化与 unroll;
  • Tensor Core fragment 的 lane 分布;
  • 对齐、临时变量和地址计算;
  • 架构的每线程与每 SM 寄存器限制。

因此应把 96 当作算法层的 live scalar 预算,不作为编译器报告的承诺。

10. 双缓冲只在能够真正并发时有效

代码里准备两个数组,不会自动让 load 与 compute 重叠。

需要同时满足:

  • 硬件和指令支持相应的异步传输;
  • producer/consumer 有正确 barrier 或完成事件;
  • 复用槽位前确认旧数据不再被读取;
  • load 时间能被足够长的计算窗口覆盖;
  • 两套缓冲没有让 occupancy 下降到抵消收益。

现代 kernel 也可能使用三阶段或更多 stage;视频用两级双缓冲讲清最小结构。

跟练与练习

原视频跟练

  • 从 00:20 标出目标 CC tile 所需的四对输入 tile。
  • 在 00:40 画出 SRAM ping 与 pong 各自能保存什么。
  • 从 01:00 写出“计算 ping / 加载 pong”的并行动作。
  • 在 02:00 从 8×328\times3232×832\times8 推导 32 轮外积。
  • 从 03:00 画出 register ping/pong 的下一轮预取。
  • 在 03:20 复算 64+2×(8+8)=9664+2\times(8+8)=96
编者练习

若一个线程改为负责 4×84\times8 输出微块,仍使用两套输入寄存器做双缓冲,按视频的标量槽位模型需要多少槽位?

查看参考答案

输出累加器需要
4×8=324\times8=32
个槽位。
每个输入 buffer 保存 4×14\times11×81\times8 两条向量,需要 4+8=124+8=12 个槽位。
两套输入 buffer 共 2424 个,因此总计
32+24=56.32+24=56.
该数字仍是 live scalar 概念预算,不保证物理寄存器分配恰好为 56。

常见误区

  • 把两级双缓冲当成复制两份完整 C:C 累加器只有一组,双份的是输入槽位。
  • 混淆两个切换粒度:SRAM 按大 tile 切换,寄存器按 KK 位置向量切换。
  • 把 ASR 的“pm/胖/碰”当成不同缓冲:画面统一为 ping/pong。
  • 认为双缓冲能消灭全部搬运时间:稳态仍由 load 与 compute 中较慢者决定。
  • 把 96 当成编译器必然分配值:它是视频中的标量槽位模型。
  • 忽略资源成本:双缓冲增加片上占用,可能降低 occupancy。

本课小结

  • SRAM ping-pong 重叠 HBM→SRAM 大 tile 搬运与当前 tile 计算。
  • 寄存器 ping-pong 重叠 SRAM→register 向量读取与当前外积计算。
  • 一个 8×88\times8 微块需要 64 个长期输出累加槽位。
  • 两套 8+88+8 输入向量槽位再占 32,视频概念预算总计 96。
  • 两级流水的单位分别是大 KK tile 与单个 kk 位置,必须分层理解。
  • 真实收益取决于异步 copy、同步、寄存器压力和 occupancy。
08

主题讲解 · 03:04

Tensor Core 如何嵌入分块 GEMM 流水线

学习目标

  • 把 Tensor Core 理解为矩阵乘加硬件路径,而不是完整 GEMM 的替代品。
  • 说明 128×128128\times128 GEMM、32×3232\times32 SRAM tile 与 16×1616\times16 子块之间的层次。
  • 区分常规 CUDA Core 外积微内核与 warp 协作的矩阵乘加接口。
  • 解释为什么一个线程不应独自保存三个完整 16×1616\times16 fragment。
  • 描述 register ping-pong 如何与 Tensor Core 计算并行。
  • 指出矩阵 shape、warp/warp-group 粒度和支持数据类型都依赖具体 GPU 架构。

前置与衔接

P51 用外积提高一个线程微块的输入复用,P52 又把 SRAM 与寄存器两级双缓冲串起来。

本课保留同一条数据流水,只把最内层计算换成硬件矩阵乘加:

D=AB+C.D=A B+C.

这类操作通常简称 MMA(matrix multiply-accumulate)。

Tensor Core 加速的是最内层高密度矩阵乘加,不会自动解决 HBM 搬运、tile 选择、同步或边界处理。

核心讲解

1. Tensor Core 是硬件矩阵乘加接口

视频把 Tensor Core 直观描述为“小型矩阵乘法计算接口”。

开发者不再把所有标量乘加逐条组织成普通 CUDA Core 指令,而是提交满足指定 shape 与数据类型的矩阵 fragment。

图 1

Tensor Core 可作为执行小矩阵乘加的硬件级接口,嵌在更大的 tiled GEMM 数据流中。

原视频 · 00:00 ↗

概念上可以写成

CfragAfragBfrag+Cfrag.C_{\text{frag}} \leftarrow A_{\text{frag}}B_{\text{frag}}+C_{\text{frag}}.

一次接口调用内部会完成许多标量乘加,但外部仍需准备输入 fragment、维护累加器并组织循环。

2. 大矩阵仍要先按 SRAM tile 流动

视频继续使用三个 128×128128\times128 矩阵,并按 T=32T=32 切成 4×44\times4 个大 tile。

目标 CC tile 需要沿 KK 维消费四对 AABB tile。

Tensor Core 并没有足够片上容量一次吞下完整 128×128128\times128 问题,因此块级分解仍不可少。

3. SRAM 级 ping-pong 仍负责隐藏 HBM 搬运

一对 32×3232\times32 输入 tile 先进入 SRAM ping。

计算当前 ping 数据时,下一对 tile 进入 pong;之后两个槽位交换。

图 2

32×32 大 tile 仍使用 SRAM ping-pong,让当前计算与下一对 HBM→SRAM 加载交叠。

原视频 · 00:40 ↗

因此 Tensor Core 只替换最内层 compute path,不替换

HBMSRAM\text{HBM}\rightarrow\text{SRAM}

的数据供应路径。

若供应速度跟不上,再强的矩阵乘加硬件也会饿死。

4. 常规 CUDA Core 路径以线程级外积更新微块

在前一课的思路里,线程从 SRAM 读取两条向量:

aRm×1,bR1×n,a\in\mathbb R^{m\times1}, \qquad b\in\mathbb R^{1\times n},

再以外积更新自己负责的 m×nm\times n 输出微块。

图 3

常规 CUDA Core 路径可由线程从 SRAM 取向量并以外积方式更新自己的输出微块。

原视频 · 01:20 ↗

ASR 把 CUDA Core 识别为“ka call”,结合画面与上下文可确定标准术语。

该路径给程序员很细的指令控制,但标量/向量乘加需要显式调度。

5. 32×3232\times32 SRAM tile 再切成 16×1616\times16 子块

视频把一个 32×3232\times32 tile 切成 2×22\times2,得到四个 16×1616\times16 子块。

每个子块由一个 warp 协作负责,而不是交给单个线程。

图 4

32×32 SRAM tile 进一步切成四个 16×16 子块,视频用一个 warp 协作负责一个子块。

原视频 · 01:40 ↗

warp 中多个 lane 分别持有 fragment 的一部分寄存器。

程序员在逻辑上操作矩阵 fragment,但数据不会完整复制到每个线程。

6. 为什么不能让单线程保存三个完整 16×1616\times16 矩阵

若把 A、B、C 三块都按标量直接放进一个线程,需要

3×16×16=7683\times16\times16=768

个标量槽位。

视频用“单线程最多 255 个寄存器”说明这显然不可行。

更稳妥的边界是:每线程寄存器上限、编码与分配粒度依架构而变,但 768 个 live scalar 远超合理的单线程预算,还会严重破坏 occupancy 或直接无法分配。

warp 协作既匹配 Tensor Core 指令的执行粒度,也把 fragment 数据分散到多个 lane。

7. Tensor Core 接口执行 fragment 级 MMA

第一轮把一组 A、B fragment 载入 register ping。

Tensor Core 执行

C(w)A(w)B(w)+C(w),C^{(w)} \leftarrow A^{(w)}B^{(w)}+C^{(w)},

并把结果累加到同一组 C fragment。

图 5

Tensor Core 对寄存器中的 A、B fragment 执行矩阵乘加,并把结果累加到 C fragment。

原视频 · 02:20 ↗

这里的“调用接口即可”表示计算细节由硬件指令承担,不表示编写 kernel 只需一个函数调用。

数据布局、对齐、fragment 加载、累加类型和同步仍必须正确。

8. Register 双缓冲继续为下一次 MMA 预取数据

Tensor Core 使用 ping 中当前 fragment 计算时,可以把下一轮 A、B fragment 从 SRAM 载入 pong。

完成事件满足后,两个槽位交换角色。

图 6

当前 fragment 的结果继续累加到同一 C 寄存器块,同时下一轮 A、B fragment 进入另一缓冲槽。

原视频 · 02:40 ↗

经过多轮 KK 方向 MMA,同一个 C fragment 持续累加。

最终各 warp 的结果组成完整 32×3232\times32 C tile,再写回 HBM。

9. 三层结构必须同时成立

层级工作对象主要目的
HBM→SRAM32×3232\times32 大 tile 对跨较慢存储层搬运并复用
SRAM→registerTensor Core 输入 fragment为下一次 MMA 预取
Tensor Core小矩阵 fragment高吞吐矩阵乘加

只有 compute 没有供数流水,Tensor Core 会等待。

只有双缓冲没有高效内层指令,也无法利用硬件峰值吞吐。

10. 16×1616\times16 与“一个 warp”是视频示意,不是永恒契约

Tensor Core 支持的 MMA shape 会随以下条件变化:

  • GPU 架构代际;
  • FP16、BF16、TF32、FP8、INT8 等输入类型;
  • 累加器类型;
  • PTX/SASS 指令族与高级库接口;
  • warp 级 MMA 或更新架构中的 warp-group MMA。

因此本课的 16×1616\times16、一个 warp 是理解画面的具体示例。

实际开发应查当前架构的官方指令和库文档,不从本视频推导所有设备的固定 fragment shape。

11. 混合精度也需要单独审视

许多 Tensor Core 路径允许较低精度输入配合较高精度累加,例如 FP16/BF16 输入、FP32 累加。

但具体支持组合与数值行为依架构和接口而定。

选择 Tensor Core 不只影响速度,还会影响舍入误差、溢出范围与可复现性。

这是编者补充的数值边界,视频主要聚焦数据流。

跟练与练习

原视频跟练

  • 从 00:00 用 D=AB+CD=AB+C 解释 Tensor Core 的接口角色。
  • 在 00:40 指出 SRAM ping/pong 没有因 Tensor Core 而消失。
  • 从 01:20 对比线程级外积与 warp 级 fragment 计算。
  • 在 01:40 复算三个 16×1616\times16 矩阵共有 768 个标量。
  • 从 02:20 标出 A/B 输入 fragment 与 C 累加 fragment。
  • 在 02:40 画出 compute ping / load pong 的寄存器时间线。
编者练习

某 Tensor Core 微内核沿 KK 方向需要 4 次 MMA。单次 fragment 加载耗时 66,MMA 耗时 1010,单位相同。 忽略同步,比较完全串行与理想 register 双缓冲的总耗时。

查看参考答案

完全串行为
4(6+10)=64.4(6+10)=64.
理想双缓冲含一次 prologue load、3 个稳态重叠阶段和最后一次 MMA:
6+3max(6,10)+10=46.6+3\max(6,10)+10 =46.
这是理想调度估算;实际还要计入 barrier、指令发射、fragment 布局和资源竞争。

常见误区

  • 认为 Tensor Core 会自动完成整个 GEMM:大矩阵分块、搬运、同步与写回仍由 kernel 组织。
  • 把 fragment 完整放在每个线程中:warp 内 lane 分担 fragment 数据。
  • 把视频的 255 当成所有 GPU 的固定上限:寄存器限制与分配规则依架构而变。
  • 认为 Tensor Core 取代了双缓冲:计算更快反而更需要稳定供数。
  • 16×1616\times16 当成所有 MMA 的唯一 shape:shape 与数据类型、架构、指令族有关。
  • 只看峰值 FLOPs:数据供应、occupancy 和数值精度同样决定可用性能。

本课小结

  • Tensor Core 是小矩阵乘加硬件路径,不是完整 GEMM 的自动实现器。
  • 128×128128\times128 问题先切成 32×3232\times32 SRAM tile,再细分成视频示意的 16×1616\times16 子块。
  • warp 协作持有 fragment,避免单线程承担不合理的寄存器压力。
  • SRAM ping-pong 隐藏 HBM 搬运,register ping-pong 隐藏下一 fragment 读取。
  • 当前 MMA 累加到同一 C fragment,下一组 A/B fragment 同时预取。
  • shape、协作粒度、数据类型和寄存器限制必须以目标 GPU 架构为准。
09

主题讲解 · 00:59

用四条任务轴直观理解 4D 并行

学习目标

  • 从二维单层计算构造 layer、context、sequence、feature 四类索引轴。
  • 解释单个上下文的计算怎样由二维堆叠为三维工作体。
  • 说明多个上下文为什么引入第四个独立任务索引。
  • 分别指出 DP、PP、CP、TP 主要切分哪条轴。
  • 用四元组定位一个设备在并行网格中的坐标。
  • 正确理解“正交”是可独立索引与组合,不是运行时完全没有依赖或通信。

前置与衔接

前面的课程已经分别出现 tensor parallel、pipeline parallel 和 context parallel。

当四种并行同时使用时,容易把名字记成一串缩写。

视频提供了一个几何模型:把训练/推理任务放进多维网格,每一种并行沿不同轴切分。

可用四个抽象索引表示:

(b,l,s,h),(b,l,s,h),

分别代表 context/microbatch、layer、sequence/token 与 hidden feature。

这些并非唯一实现布局,却能帮助定位 DP、PP、CP、TP 的基本分工。

核心讲解

1. “正交”先从独立坐标轴理解

几何中,正交向量互相垂直。

4D 并行借用这个词,强调四种切分作用在不同任务索引上。

若固定另外三项,只改变其中一个坐标,就沿一条独立轴移动。

工程上常写总设备数为

N=NDPNPPNCPNTP,N =N_{\text{DP}} N_{\text{PP}} N_{\text{CP}} N_{\text{TP}},

前提是完整采用这四维笛卡尔网格且每张设备只对应一个坐标元组。

2. 单层计算可先画成二维平面

对一个 Transformer layer,可以把激活直观画成

XlRS×H,X_l\in\mathbb R^{S\times H},

其中 SS 是 token/sequence 轴,HH 是 hidden feature 轴。

attention 与 FFN 都在这张二维激活上产生新的二维输出。

画面下方的小型 Q/K/A/V/O/FFN 流程只是单层二维计算的示意,不是完整 tensor shape 规范。

3. 沿 layer 轴堆叠得到单上下文三维工作体

把第 1 层到第 LL 层的二维激活依次叠起来,得到

XRL×S×H.X\in\mathbb R^{L\times S\times H}.
图 1

把单层二维 token×feature 计算沿 layer 轴堆叠,可形成单个上下文的三维工作体。

原视频 · 00:20 ↗

这里的三维不是说运行时必须保存整个 L×S×HL\times S\times H 张量。

它是任务坐标的可视化:每个 layer 都有自己的 token×feature 计算平面。

4. 多个上下文引入第四个轴

单个三维工作体只表示一个上下文或一个 microbatch 项。

多个相互独立的上下文可以写成

XRB×L×S×H.X\in\mathbb R^{B\times L\times S\times H}.
图 2

多个独立上下文各自形成一个三维工作体,把 context/batch 作为额外轴后得到四维任务空间。

原视频 · 00:30 ↗

BB 在这里是任务级 context/microbatch 轴,不应机械等同于某个框架张量中的唯一 batch 维。

数据加载、gradient accumulation 和 pipeline microbatch 都可能改变它的物理组织。

5. DP 沿独立样本或 microbatch 轴切分

数据并行把不同输入样本或 microbatch 分给不同 replica。

图 3

数据并行 DP 沿独立样本或 microbatch 轴分配不同上下文工作体。

原视频 · 00:40 ↗

最基础的 DP 在每个 replica 上保留完整模型参数,分别前向/反向,再同步梯度。

FSDP/ZeRO 会进一步分片参数、梯度或优化器状态,但它们仍服务于数据并行语义。

6. PP 沿 layer 轴切分

流水线并行把连续或规则分组的网络层放到不同 stage。

图 4

流水线并行 PP 沿网络层轴把连续层段放到不同 stage。

原视频 · 00:45 ↗

相邻 stage 之间传递激活与梯度。

为了提高利用率,microbatch 会在各 stage 间流水,但仍存在 pipeline bubble、调度与负载均衡问题。

7. CP 沿长上下文的 sequence 轴切分

上下文并行把单个长序列的 token 工作分给多张设备。

图 5

上下文并行 CP 沿单个长上下文的 sequence/token 轴切成多个连续或布局相关的片段。

原视频 · 00:50 ↗

最简单可想成连续 token 段,但实际实现也可能采用 ring、blockwise 或其他布局。

attention 需要跨 token 汇总 K/V 或中间统计量,所以 CP 不是“切完完全不通信”。

8. TP 沿特征或算子内部维度切分

张量并行在单层内部切分权重与激活。

常见切分对象包括 hidden feature、attention head、FFN intermediate feature 或矩阵乘法的输入/输出维。

图 6

张量并行 TP 沿隐藏特征、attention head 或算子内部相关维度切分单层计算。

原视频 · 00:55 ↗

因此“TP=特征维”是视频的直觉入口,不是说所有 TP 都固定切同一个物理轴。

列并行、行并行和 head parallel 都属于具体实现。

9. 一个设备由四维坐标定位

(NDP,NPP,NCP,NTP)=(2,4,2,8),(N_{\text{DP}},N_{\text{PP}},N_{\text{CP}},N_{\text{TP}}) =(2,4,2,8),

则完整并行网格需要

2×4×2×8=1282\times4\times2\times8=128

个 rank。

每个 rank 可由

(iDP,iPP,iCP,iTP)(i_{\text{DP}},i_{\text{PP}},i_{\text{CP}},i_{\text{TP}})

唯一定位。

固定其他三项、只改变一项,就得到对应维度的通信组。

例如 TP group 固定 DP/PP/CP 坐标,只遍历 iTPi_{\text{TP}}

10. 四维“正交”不等于运行时完全解耦

四个轴可分别索引和组合,但真实系统仍有耦合:

  • TP collective 与 CP attention 通信可能争用同一网络;
  • PP stage 的计算量会受 TP/CP 切分影响;
  • DP 梯度同步时机依赖 pipeline schedule;
  • 并行度要匹配 head 数、层数、序列长度和硬件拓扑;
  • 不同维度的通信组可能跨节点或限于节点内。

所以“正交”是任务分解的结构性质,不是性能独立性或零通信承诺。

11. 4D 是并行维度数量,不是张量 rank 的硬性要求

模型中的某个激活张量可能有 batch、sequence、head、head-dim 等更多轴,也可能在实现中被展平。

“4D parallelism”说的是四类并行策略的组合,不等同于所有中间张量必须是 rank-4。

视频用长方体帮助直观理解,不能据此替代实际 shape 与通信组设计。

跟练与练习

原视频跟练

  • 从 00:20 把二维 token×feature 平面沿 layer 轴堆叠。
  • 在 00:30 说明多个 context 为什么是新的独立索引。
  • 从 00:40 指出 DP 切的是哪个工作体集合。
  • 在 00:45 指出 PP stage 与 layer 段的对应。
  • 从 00:50 标出 CP 的 sequence/token 轴。
  • 在 00:55 列举 TP 可切分的 feature/head/矩阵维度。
编者练习

一个训练作业配置为 DP=4、PP=2、CP=2、TP=8。 总共需要多少个 rank?某个 TP group 内哪些坐标固定,哪个坐标变化?

查看参考答案

若采用完整四维笛卡尔网格,总 rank 数为
4×2×2×8=128.4\times2\times2\times8=128.
一个 TP group 固定
(iDP,iPP,iCP),(i_{\text{DP}},i_{\text{PP}},i_{\text{CP}}),
只让
iTP=0,1,,7i_{\text{TP}}=0,1,\ldots,7
变化,因此该 group 有 8 个 rank。

常见误区

  • 把正交理解为四种并行完全不通信:它表示切分轴可独立索引,不表示运行时无耦合。
  • 把 4D 等同于激活张量必须 rank-4:这里描述四类并行策略组合。
  • 把 DP 只理解为复制完整参数:FSDP/ZeRO 可进一步分片状态。
  • 把 PP 当成切 token:PP 的主要轴是网络层/stage。
  • 把 CP 当成普通 batch 切分:CP 切单个长上下文的 token 工作。
  • 把 TP 固定为一个唯一特征方向:具体可切 hidden、head、FFN 或矩阵维。

本课小结

  • 单层 token×feature 平面沿 layer 堆叠形成单上下文三维任务体。
  • 多个 context/microbatch 工作体再引入第四个独立轴。
  • DP、PP、CP、TP 分别主要对应 context、layer、sequence、feature/算子内部轴。
  • 完整并行网格的 rank 数可写为四个并行度的乘积。
  • 每张设备由四元组坐标定位,对应通信组通过固定三轴、遍历一轴构造。
  • “正交”是任务索引抽象;通信、调度、负载与硬件拓扑仍会互相影响。
10

主题讲解 · 03:38

Megatron 如何让注意力按头分片后只在出口通信

学习目标

  • 用 shape 推导线性层列切分与行切分的不同合并语义。
  • 解释标准多头注意力为什么天然能按 head 分给不同 GPU。
  • 推导 WQ,WK,WVW_Q,W_K,W_V 列切分后的局部 Q/K/V shape。
  • 说明每张卡为何能独立完成本地 head 的 QKQK^\top、softmax 与 AVAV
  • 推导 WOW_O 行切分后的局部部分和与最终 All-Reduce。
  • 指出 GQA/MQA、head 整除、反向通信、sequence parallel 和 bias 的实现边界。

前置与衔接

P47 已建立 tensor parallel 的两条基本规则:

  • column parallel 产生不同输出特征分片,数学合并是 concat;
  • row parallel 产生同形状部分和,数学合并是 sum。

P50 把两者配成 FFN 的“先列后行”,让逐元素激活在中间分片上直接执行。

本课把同一思想迁移到多头自注意力:

column-parallel QKVlocal attention headsrow-parallel WO.\text{column-parallel QKV} \rightarrow \text{local attention heads} \rightarrow \text{row-parallel }W_O.

中间不拼接全部 head,出口再做一次求和通信。

核心讲解

1. 先统一线性层 shape

把 batch 与 sequence 等非隐藏维展平为 MM,写成

XRM×K,WRK×N,Y=XWRM×N.X\in\mathbb R^{M\times K}, \qquad W\in\mathbb R^{K\times N}, \qquad Y=XW\in\mathbb R^{M\times N}.

判断行列并行时,以数学公式中的 WW 为准。

若实现保存 WW^\top,内存画面中的行列会反转,但收缩维和输出维的语义不会变。

2. 列切分沿输出特征分配工作

WW 沿列分为

W=[W0  W1    Wp1],WrRK×N/p.W=[W_0\;W_1\;\cdots\;W_{p-1}], \qquad W_r\in\mathbb R^{K\times N/p}.
图 1

线性层的列切分复制输入 X,并让各卡持有 W 的不同输出列与对应输出特征分片。

原视频 · 00:20 ↗

每张卡使用完整输入计算

Yr=XWrRM×N/p.Y_r=XW_r\in\mathbb R^{M\times N/p}.

完整结果逻辑上为 Y=[Y0;Y1;]Y=[Y_0;Y_1;\ldots]

若下游能直接消费这些输出分片,就无需立即 all-gather。

3. 行切分沿收缩维产生部分和

WW 沿行分为

W=[W0W1Wp1],WrRK/p×N.W= \begin{bmatrix} W_0\\W_1\\\vdots\\W_{p-1} \end{bmatrix}, \qquad W_r\in\mathbb R^{K/p\times N}.

输入也按相同 KK 维分片:

X=[X0  X1    Xp1].X=[X_0\;X_1\;\cdots\;X_{p-1}].

局部结果

Pr=XrWrRM×NP_r=X_rW_r\in\mathbb R^{M\times N}

都是完整输出 shape,但只包含一段 KK 维贡献。

图 2

线性层的行切分同时切分 X 的输入特征,各卡产生同形状部分和,最终通过 sum All-Reduce 合并。

原视频 · 01:00 ↗

因此

Y=r=0p1Pr.Y=\sum_{r=0}^{p-1}P_r.

若每卡都需要完整 YY,使用 sum All-Reduce。

4. 标准多头注意力先按 head 拆成独立分支

设标准 MHA 有 hh 个 head,每个 head 宽度为 dhd_h,并令

dmodel=hdh.d_{\text{model}}=h d_h.

输入

XRM×dmodel.X\in\mathbb R^{M\times d_{\text{model}}}.

投影得到

Q=XWQ,K=XWK,V=XWV,Q=XW_Q, \quad K=XW_K, \quad V=XW_V,

其中在标准 MHA 中

WQ,WK,WVRdmodel×hdh.W_Q,W_K,W_V \in\mathbb R^{d_{\text{model}}\times h d_h}.

输出特征可以按 head 分组,正适合列切分。

5. Q/K/V 投影沿 head 对应的输出列切分

先看视频的两头两卡示例。

对每个投影矩阵写成

WQ=[WQ,0  WQ,1],W_Q=[W_{Q,0}\;W_{Q,1}],
WQ,rRdmodel×dh,W_{Q,r}\in\mathbb R^{d_{\text{model}}\times d_h},

WK,WVW_K,W_V 同理。

图 3

WQ、WK、WV 沿输出特征列切分后,各卡直接得到自己负责 attention head 的 Q、K、V。

原视频 · 01:40 ↗

两张卡都使用完整 XX,分别得到

Qr=XWQ,r,Kr=XWK,r,Vr=XWV,r,Q_r=XW_{Q,r}, \quad K_r=XW_{K,r}, \quad V_r=XW_{V,r},

Qr,Kr,VrRM×dh.Q_r,K_r,V_r\in\mathbb R^{M\times d_h}.

每卡已经拥有自己 head 的完整 Q/K/V,无需为投影结果立即通信。

6. 每个标准 attention head 可在本地独立计算

rr 个 head 执行

Ar=softmax(QrKrdh+Mr),A_r =\operatorname{softmax} \left( \frac{Q_rK_r^\top}{\sqrt{d_h}}+M_r \right),
Or=ArVr.O_r=A_rV_r.

这里 MrM_r 表示相应的 mask;标准情形下各 head 使用兼容的 mask 规则。

图 4

每张卡在本地完成自己 head 的 QKᵀ、softmax 与 AV,头间在标准 MHA 核心计算中互不混合。

原视频 · 02:00 ↗

标准 MHA 的核心 attention 不在 head 之间混合数值。

因此 O0,O1O_0,O_1 暂不 concat,继续保留在各卡即可。

“头之间互不干涉”只针对这一段核心计算;最终输出投影正是重新混合 head 信息的位置。

7. WOW_O 行切分正好接住本地 head 输出

完整 head 输出逻辑上为

Ocat=[O0  O1    Oh1]RM×hdh.O_{\text{cat}}=[O_0\;O_1\;\cdots\;O_{h-1}] \in\mathbb R^{M\times h d_h}.

输出投影

WORhdh×dmodel.W_O\in\mathbb R^{h d_h\times d_{\text{model}}}.

按负责的 head 组沿 WOW_O 行切分:

WO=[WO,0WO,1WO,p1].W_O= \begin{bmatrix} W_{O,0}\\W_{O,1}\\\vdots\\W_{O,p-1} \end{bmatrix}.

若每卡负责 h/ph/p 个 head,则

WO,rR(h/p)dh×dmodel.W_{O,r} \in \mathbb R^{(h/p)d_h\times d_{\text{model}}}.

本卡局部 head 输出拼成

OrlocalRM×(h/p)dh,O_r^{\text{local}} \in \mathbb R^{M\times(h/p)d_h},

收缩维恰好对齐。

8. 各卡先得到 dmodeld_{\text{model}} 宽的部分和

rr 张卡计算

Pr=OrlocalWO,rRM×dmodel.P_r =O_r^{\text{local}}W_{O,r} \in \mathbb R^{M\times d_{\text{model}}}.

每个 PrP_r shape 都已经恢复到模型宽度,但只含本卡 head 组对最终输出的贡献。

完整输出为

Y=r=0p1Pr.Y=\sum_{r=0}^{p-1}P_r.
图 5

局部 head 输出直接进入行切分的 WO,各卡得到 d_model 宽的部分和,再执行一次 sum All-Reduce。

原视频 · 03:00 ↗

若所有卡都需要复制的 YY,执行一次 sum All-Reduce。

这就是注意力前向中“入口列切分、中间本地算、出口行切分”的通信压缩路径。

9. 16 个 head 与 4 张 GPU 的推广

h=16,p=4h=16,p=4,每卡负责 4 个 head。

本卡 Q/K/V 输出宽度为

4dh.4d_h.

本地完成四个 head 后,临时 concat 为

OrlocalRM×4dh.O_r^{\text{local}} \in\mathbb R^{M\times4d_h}.

对应的 WO,rW_{O,r} 行分片为

WO,rR4dh×dmodel.W_{O,r} \in\mathbb R^{4d_h\times d_{\text{model}}}.
图 6

当 16 个 attention head 分给 4 张 GPU 时,每卡可负责 4 个 head,再接 WO 的对应四分之一行分片。

原视频 · 03:20 ↗

四张卡各自产生 M×dmodelM\times d_{\text{model}} 部分和,最终一次 All-Reduce。

10. 省掉的是中间 head all-gather

若在 Q/K/V 投影后立即拼接所有 head,再重新分给后续算子,会产生无谓通信。

Megatron 风格布局让:

  1. Q/K/V 列并行输出直接保留为 head 分片;
  2. attention 核心在本卡消费这些分片;
  3. WOW_O 行分片直接消费本地 head 输出;
  4. 只有形成最终模型宽度输出时求和。

优化来自上下游布局连续性,不是注意力公式发生变化。

11. “一次通信”限于视频的简化前向主干

本课的一次 All-Reduce 指标准 MHA 这段前向路径的主要同步点。

训练反向还需要为输入梯度、权重梯度和分片布局执行相应 collective。

若使用 sequence parallel,出口也可能采用 reduce-scatter,让结果继续保持序列分片,而不是复制完整输出。

因此不能由本图推出完整训练 step 只通信一次。

12. GQA/MQA 不能机械套“Q/K/V 同样按 head 切”

标准 MHA 通常有相同数量的 Q、K、V head。

GQA 让多个 Q head 共享较少的 KV head,MQA 甚至可能只有一组 KV head。

此时 WQW_QWK,WVW_K,W_V 的输出宽度和分组不同。

实现可能:

  • 分片 Q head、复制 KV head;
  • 按 KV group 对齐 TP rank;
  • 在特定阶段执行 KV 通信;
  • 限制 TP 度与 KV head 数的整除关系。

所以视频中的“Q/K/V 一样”只对标准 MHA 两头示例成立。

13. Head 数、bias 与物理权重布局也有边界

最简单的 head parallel 要求

hmodp=0.h\bmod p=0.

若不能整除,需要不均匀分配、复制、padding 或换并行度。

很多实现把 WQ,WK,WVW_Q,W_K,W_V 融合成一个大权重,物理切片顺序依框架而异;数学上仍应追踪每个 rank 拥有哪些输出 head/features。

WOW_O 带 bias,不能让每个部分和都加一份完整 bias 后再求和,否则会重复 pp 次。

常见做法是归约后加一次,或采用保证全局只贡献一次的约定。

跟练与练习

原视频跟练

  • 从 00:20 判断列切分的局部结果为何应 concat。
  • 在 01:00 判断行切分的局部结果为何应 sum。
  • 从 01:40 给两头两卡写出 WQ,r:dmodel×dhW_{Q,r}:d_{\text{model}}\times d_h
  • 在 02:00 逐卡写出 Ar=softmax(QrKr/dh)A_r=\operatorname{softmax}(Q_rK_r^\top/\sqrt{d_h})
  • 从 03:00 检查 OrlocalWO,rO_r^{\text{local}}W_{O,r} 的收缩维。
  • 在 03:20 把 16 个 head 均分给 4 张卡并复算局部宽度。
编者练习

标准 MHA 有 dmodel=4096d_{\text{model}}=4096h=32h=32,因此 dh=128d_h=128。使用 TP=8。 写出每卡负责的 head 数、局部 Q 输出宽度、WO,rW_{O,r} shape 与局部输出部分和 shape;把非隐藏维记为 MM

查看参考答案

每卡负责
32/8=432/8=4
个 head。
局部 Q/K/V 宽度为
4×128=512,4\times128=512,
所以局部 Q shape 为 M×512M\times512
WOW_O 行分片为
WO,rR512×4096.W_{O,r}\in\mathbb R^{512\times4096}.
局部乘积
PrRM×4096.P_r \in\mathbb R^{M\times4096}.
八个 PrP_r 是同形状部分和,最终用 sum All-Reduce 合并。

常见误区

  • 把列切分的输出直接相加:不同 rank 拥有不同 head/features,数学上是 concat 分片。
  • 在 Q/K/V 后立即 all-gather:标准 MHA 的本地 head 可以直接继续计算。
  • WOW_O 继续按列切:视频方案需要按输入/head 维行切以接住本地输出。
  • 把出口部分和 concat:每个局部结果都已是 dmodeld_{\text{model}} 宽,应求和。
  • 认为头之间永远互不影响WOW_O 会混合所有 head 的贡献。
  • 把标准 MHA 的 K/V 分片机械套到 GQA/MQA:KV head 数和共享关系不同。
  • 声称训练全程只有一次通信:结论仅限简化前向主干。

本课小结

  • Q/K/V 投影按输出列切分后,每卡直接得到自己负责的 attention head。
  • 标准 MHA 的 QKQK^\top、softmax 与 AVAV 可在各 head 内本地完成。
  • 本地 head 输出不必全局 concat,可直接进入 WOW_O 的对应行分片。
  • 每卡产生 M×dmodelM\times d_{\text{model}} 部分和,出口用 sum All-Reduce 合并。
  • 16 头、4 GPU 时每卡负责 4 头及 WOW_O 的四分之一输入行。
  • GQA/MQA、head 整除、backward、sequence parallel、bias 与融合权重都需按实现重新核对。
11

单元综合

从 GEMM 数据复用到 4D 并行:计算、通信与布局的统一账本

单元能力目标

完成本单元后,应能沿着“单卡微内核—片上流水—多卡分片—集体通信—多维并行”逐层分析一个 Transformer 线性层。

具体需要做到:

  • C=ABC=AB 推导单个输出、外积更新和 tiled GEMM;
  • 区分 HBM、shared memory/SRAM 与 register 的复用层级;
  • 说明 ping-pong 双缓冲隐藏什么 latency、增加什么资源压力;
  • 判断 Tensor Core 位于 GEMM 流水线的哪一级;
  • 从 shape 区分 column parallel 与 row parallel;
  • 计算 Ring All-Reduce 的单 rank 通信量并统一收发口径;
  • 解释 Megatron FFN 与 attention 为什么可把通信推迟到出口;
  • 用四元组坐标理解 DP、PP、CP、TP 的并行组构造与相互影响。

概念连接

1. GEMM 的数学工作不因分块减少

完整 GEMM 常写为

Cout=αAB+βCin.C_{out}=\alpha AB+\beta C_{in}.

ARM×K,BRK×N,A\in\mathbb R^{M\times K}, \qquad B\in\mathbb R^{K\times N},

CRM×N.C\in\mathbb R^{M\times N}.

单个输出为

Cij=k=1KAikBkj.C_{ij} = \sum_{k=1}^{K} A_{ik}B_{kj}.

分块不改变 O(MNK)O(MNK) 乘加次数,核心是提高输入数据在片上的复用,减少慢速内存流量。

2. tiled GEMM 沿三个轴分块

对输出 tile

CI,J,C_{I,J},

沿收缩维依次读取

AI,KtA_{I,K_t}

BKt,J.B_{K_t,J}.

每对输入 tile 被线程块内多个输出元素共同使用。

理想方阵模型中,边长为 TT 的 tile 可让输入 HBM 流量相对朴素 O(N3)O(N^3) 读取下降到约

O(N3/T).O(N^3/T).

实际收益受 cache、tile 边界、同步、bank conflict 和 occupancy 影响。

3. shared memory 与 register 承担不同复用

shared memory 或片上 SRAM 保存线程块级 A/B tile,让多个线程复用。

register 保存每个线程或 warp 的细粒度输入 fragment 与输出累加器。

典型层次为:

HBMSRAM/sharedregisterMMA/FMA.HBM \rightarrow SRAM/shared \rightarrow register \rightarrow MMA/FMA.

把所有片上存储统称“cache”会丢失显式搬运、同步和容量差异。

4. 外积适合更新一个输出微块

矩阵乘法也可写成外积和:

AB=k=1KA:,kBk,:.AB = \sum_{k=1}^{K} A_{:,k}B_{k,:}.

4×44\times4 输出微块,一个 4×14\times1 列向量与 1×41\times4 行向量的外积一次产生 16 个乘积。

若保留 16 个长期累加器,只需加载 4 个 A 元素与 4 个 B 元素,就能更新整块。

视频概念模型的 live scalar 账本为:

  • 8 个当前输入槽位;
  • 16 个输出累加槽位;
  • 合计 24。

它说明逻辑复用,不是直接可移植的硬件寄存器分配。

5. 微块越大,复用与寄存器压力同时上升

更大的输出微块能让每个输入元素参与更多乘加。

但累加器数量按面积增长,可能导致:

  • register pressure 上升;
  • occupancy 下降;
  • spill 到 local memory;
  • 线程或 warp 协作更复杂。

所以微块尺寸需要结合 ISA、数据类型、寄存器文件和目标架构选择。

6. 第一层 ping-pong:HBM 到 SRAM

沿 KK 维处理当前 tile 时,可让两组 SRAM buffer 交替工作:

  • compute 使用 buffer A;
  • async copy 把下一对 A/B tile 填入 buffer B;
  • 同步后交换角色。

单缓冲阶段耗时近似

TL+TC.T_L+T_C.

理想稳态双缓冲近似

max(TL,TC),\max(T_L,T_C),

但还存在 fill/drain 与同步开销。

7. 第二层 ping-pong:SRAM 到 register

在当前外积或 MMA 使用一组 register fragment 时,可预取下一组 fragment。

两级缓冲的单位不同:

  • 外层以大 KK tile 为单位隐藏 HBM 搬运;
  • 内层以单个 kk 步或 MMA fragment 为单位隐藏 SRAM 读取。

视频 8×88\times8 微块的概念预算包含:

  • 64 个输出累加槽位;
  • 两套各 8+88+8 的输入向量槽位,共 32;
  • 合计 96 个 live scalar。

真实 register allocation 必须看编译器与目标架构。

8. 双缓冲不减少 bytes 或 FLOPs

它用更多片上空间与更复杂同步换取 latency overlap。

完全重叠要求:

  • 异步 copy 支持;
  • 搬运与计算资源可并发;
  • 数据依赖与 barrier 正确;
  • 当前计算足够长;
  • tile 与缓冲没有把 occupancy 压垮。

若 compute 很短或带宽已经饱和,双缓冲收益可能有限。

9. Tensor Core 是微内核硬件,不是完整 GEMM

Tensor Core / MMA 指令对小矩阵 fragment 执行乘加。

大型矩阵仍需:

  1. 分块;
  2. HBM 到 SRAM 搬运;
  3. warp 协作加载 fragment;
  4. 多次 MMA 累加到 C fragment;
  5. 写回结果;
  6. 处理边界与布局。

视频中的 16×1616\times16 子块是教学示意;支持的 MMA shape、数据类型和协作粒度依 GPU 架构而变。

10. column parallel 切输出维

Y=XW,Y=XW,

XRM×K,WRK×N.X\in\mathbb R^{M\times K}, \qquad W\in\mathbb R^{K\times N}.

column parallel 将 WW 沿输出维 NN 切分:

W=[W1W2Wp].W= \begin{bmatrix} W_1&W_2&\cdots&W_p \end{bmatrix}.

每卡复制或接收完整 XX,计算

Yr=XWrRM×N/p.Y_r=XW_r \in\mathbb R^{M\times N/p}.

数学合并是沿输出维 concat。

是否立刻 All-Gather 取决于下游能否继续消费分片布局。

11. row parallel 切收缩维

row parallel 沿共享维 KK 切分:

X=[X1X2Xp],X= \begin{bmatrix} X_1&X_2&\cdots&X_p \end{bmatrix},
W=[W1W2Wp].W= \begin{bmatrix} W_1\\W_2\\\vdots\\W_p \end{bmatrix}.

每卡计算部分和

Zr=XrWrRM×N.Z_r=X_rW_r \in\mathbb R^{M\times N}.

最终

Y=r=1pZr.Y=\sum_{r=1}^{p}Z_r.

所以数学合并是求和,常用 All-Reduce 或 reduce-scatter。

12. 权重存储转置会制造“行列错觉”

PyTorch 常把线性层权重存为 [out,in],而数学公式常写 XWXWWW[in,out]

因此判断 column/row parallel 时应先写:

  • 输入 shape;
  • 数学乘法方向;
  • 收缩轴;
  • 输出轴;
  • 本地分片 shape。

不要根据代码张量的视觉行列直接命名。

13. Megatron FFN 用先列后行保持分片连续

对标准 FFN:

H=GELU(XWup),H=\operatorname{GELU}(XW_{up}),
Y=HWdown.Y=HW_{down}.

上投影使用 column parallel,得到每卡

HrRM×4D/p.H_r\in\mathbb R^{M\times4D/p}.

GELU 逐元素局部执行,不改变分片布局。

下投影使用 row parallel,使 HrH_rWdown,rW_{down,r} 在收缩维对齐,每卡得到完整 shape 的部分和。

只在出口做一次 sum collective,避免在中间 all-gather 全部 4D4D 激活。

14. attention 也可按 head 保持局部性

标准 MHA 中,Q/K/V 投影按输出列切分后,每卡直接拥有一组完整 head。

单个 head 内的

QKT,softmax,AVQK^T, \quad softmax, \quad AV

都可本地计算。

本地 head 输出直接进入输出投影 WOW_O 的对应 row 分片。

每卡产生完整输出 shape 的部分和,出口再求和。

GQA/MQA、head 数不整除、sequence parallel 与 fused QKV 都需按实际实现重新核对。

15. Ring All-Reduce 拆成两个阶段

张量大小为 SS,rank 数为 NN

Ring All-Reduce 把张量分成 NN 个 chunk:

  1. reduce-scatter:每卡最终保留一个已归约 chunk;
  2. all-gather:分发所有已归约 chunk。

每阶段有 N1N-1 步,每步发送 S/NS/N

单 rank 总发送量为

2(N1)SN=2(11N)S.2(N-1)\frac{S}{N} = 2\left(1-\frac1N\right)S.

NN 增大时趋近 2S2S

16. 通信量趋近常数不等于延迟恒定

若统计收发合计,需要再乘二。

比较数字前必须统一:

  • 单向发送量;
  • 接收量;
  • 收发合计;
  • payload 还是链路总流量。

Ring 的步数随 NN 增长,实际延迟还受拓扑、链路、chunk、启动成本和竞争影响。

17. 4D 并行是任务索引轴的组合

课程用四条主要任务轴理解:

  • DP:不同 data/context 或 microbatch;
  • PP:不同 layer stage;
  • CP:长 sequence/context 分片;
  • TP:feature、head 或算子内部张量分片。

总 rank 数可写为

Nrank=dDPdPPdCPdTP.N_{rank} = d_{DP}d_{PP}d_{CP}d_{TP}.

每张设备由四元组坐标定位。

构造某一并行组时,固定另外三轴,只遍历目标轴。

18. “正交”不等于性能互不影响

四轴在索引空间可独立组合,但真实系统仍共享:

  • 网络拓扑与带宽;
  • GPU 内存;
  • pipeline bubble;
  • batch 与 microbatch;
  • kernel shape;
  • 通信调度。

所以并行配置是联合优化问题,不能分别选择每轴最大值后直接相乘。

对比与决策

1. 单卡优化的顺序

  1. 确定 GEMM shape 与数据类型。
  2. 用 tile 提高 HBM 到 SRAM 复用。
  3. 用微块与 register 累加提高细粒度复用。
  4. 在硬件支持下映射到 MMA/Tensor Core。
  5. 用双缓冲重叠搬运,并检查 occupancy。
  6. 以 profiler 验证吞吐、带宽、stall 与 register spill。

2. 选择 column 还是 row parallel

  • 下游可继续使用输出分片:优先 column parallel,延迟 concat/all-gather。
  • 输入已经沿收缩维分片:row parallel 可直接产生部分和。
  • 两层连续线性层:尝试先 column 后 row,让中间逐元素算子本地执行,只在出口通信。

3. collective 选择取决于下游布局

  • 所有 rank 都需要完整求和结果:All-Reduce。
  • 下游只需结果分片:reduce-scatter。
  • 已有不同输出分片,需要拼完整:All-Gather。

数学上是 concat 或 sum,只决定合并语义;具体 collective 还取决于下一层布局。

综合训练

编者练习

XR32×4096X\in\mathbb R^{32\times4096}WR4096×16384W\in\mathbb R^{4096\times16384} 使用 4 卡 column parallel。写出每卡权重和输出 shape;若下一层沿 16384 输入维做 row parallel,说明为何中间不必 All-Gather。

查看参考答案

每卡 WrR4096×4096W_r\in\mathbb R^{4096\times4096},输出 Hr=XWrR32×4096H_r=XW_r\in\mathbb R^{32\times4096}。下一层若沿其输入收缩维 16384 切成 4 份,则每卡正好需要自己的 HrH_r 与对应权重行分片,能本地计算完整输出 shape 的部分和。中间逐元素激活也可在本地分片上执行,只在第二层出口对部分和做 All-Reduce 或 reduce-scatter。

编者练习 2

一个 Ring All-Reduce 张量大小为 1 GiB,使用 8 卡。计算单卡发送量;若报告收发合计,应是多少?

查看参考答案

单卡发送量为 2(11/8)S=1.752(1-1/8)S=1.75 GiB。对称 ring 中接收量同样约 1.75 GiB,收发合计约 3.5 GiB。必须说明统计口径;“单卡通信量 1.75 GiB”通常指发送 payload,而不是收发总字节。

编者练习 3

某双缓冲 GEMM 的单 tile 加载为 6 微秒,计算为 10 微秒,共 20 个 tile。比较理想单缓冲和忽略额外同步时的双缓冲时间,并指出现实中还需检查什么。

查看参考答案

单缓冲约为 20(6+10)=32020(6+10)=320 微秒。理想双缓冲包含首个 load、稳态重叠和排空,近似 6+20×10=2066+20\times10=206 微秒,具体边界计数依实现。现实中还需检查异步 copy、barrier、buffer 容量、register/shared memory 压力、occupancy、带宽竞争和 tile 尾部。

进入下一单元前

  • 已能从内积和外积两种视角解释 GEMM 微内核。
  • 已能分开 HBM→SRAM 与 SRAM→register 两层缓冲。
  • 已能说明 Tensor Core 只是分块流水中的小矩阵乘加路径。
  • 已能从 shape 区分 column/row parallel 与 concat/sum 合并。
  • 已能计算 Ring All-Reduce 的单 rank 发送量并统一收发口径。
  • 已能解释 FFN 与 MHA 如何保持中间分片并把通信推迟到出口。
  • 已能用四元组坐标构造 DP/PP/CP/TP 组,并说明它们在性能上不独立。
  • 若仍会由数学转置误判行列并行,回看 P47。
  • 若仍把 2S2S 解释为恒定延迟,回看 P48。
  • 若仍把 Tensor Core 当作完整 GEMM,回看 P53 的数据流水账本。