CS336 学习笔记 02:注意力替代、MoE 与 GPU 系统

本文最后更新于 2026年9月28日 晚上

CS336 入门辅助学习文档 — 系统与并行模块


写在前面

前三讲我们了解了大模型的"骨架"——Transformer 架构、分词、资源核算。从这一讲开始,我们进入系统与并行模块,要回答一个核心问题:

这么大的模型,一张 GPU 装不下、训不动,怎么办?

这五讲会带你从"一张 GPU 怎么跑得快"(L5-L6),到"多张 GPU 怎么一起干活"(L7-L8),再到"模型本身能不能拆得更高效"(L4)。

前置知识

在开始之前,你需要知道: - 什么是 Transformer(第三讲内容) - 什么是注意力机制(第三讲内容) - GPU 是什么,显存和计算单元的区别(第二讲内容) - 什么是 FLOPs 和内存带宽(第二讲内容)

如果这些概念你还不太清楚,建议先回去看看前三讲的辅助文档。


第四讲:注意力替代方案与混合专家模型(MoE)

4.0 前置知识:注意力的"痛点"是什么?

在第三讲我们学过,注意力机制是 Transformer 的核心。它让每个 token 都能"看到"其他所有 token,就像开会时每个人都能和所有人交流。

但注意力有个大问题:它的计算量是 O(n²) 的。

什么意思呢?想象一下: - 序列长度是 100 个 token → 注意力需要 100×100 = 10,000 次计算 - 序列长度是 1000 个 token → 注意力需要 1000×1000 = 1,000,000 次计算 - 序列长度是 10,000 个 token → 注意力需要 10,000×10,000 = 100,000,000 次计算

序列长度翻 10 倍,计算量翻 100 倍!这就是"二次方增长"。

短文本还好,但现在大模型的上下文窗口越来越大——从 4K 到 32K 到 128K 到 1M(100万)。如果还用标准注意力,计算量会爆炸。

💡 生活类比: 想象一个 100 人的会议,每个人都要和其他人一对一交流。总共需要 100×99/2 ≈ 5000 次交流。 如果是 1000 人的会议呢?需要约 50 万次交流!人越多,交流成本增长越快。 注意力机制就是这样——序列越长,计算开销越大。

所以,人们开始想:有没有办法让注意力变得更高效?

这就是这一讲要讲的:注意力的各种替代方案,以及另一种思路——混合专家模型(MoE)。


4.1 线性注意力:把二次方变成线性

4.1.1 核心思想:换个计算顺序

标准注意力的公式是这样的:

1
Attention(Q, K, V) = softmax(QK^T / √d_k) V

先算 Q 和 K 的内积(得到一个 n×n 的大矩阵),再乘 V。这个 n×n 的大矩阵就是 O(n²) 的来源。

注意公式里要除以 √d_k(每个头的维度)做缩放:维度一大,点积的数值会随 √d_k 增大,softmax 容易饱和成“一家独大”、梯度几乎消失;除以 √d_k 把点积方差拉回 1,训练才稳定。

那能不能换个顺序,先算 K 和 V,再乘 Q 呢?

1
Q (K^T V)

这样就变成了: - 先算 K^T V:一个 d×n 的矩阵乘一个 n×d 的矩阵,得到 d×d 的小矩阵(d 是维度,通常几百) - 再算 Q 乘这个 d×d 的小矩阵:n×d 乘 d×d,得到 n×d 的结果

整个过程的计算量是 O(n × d²),而 d 是固定的(比如 512),所以相对于 n 来说是线性的!

💡 生活类比: 想象你要计算"每个学生的各科成绩加权总分"。 - 方法一(标准注意力):先算每个学生和每门课的"相关性"(一个学生×课程的大表),再用这个大表去乘课程的分数。 - 方法二(线性注意力):先把课程分数按权重合并成一个"综合分"(课程数×维度的小表),再让每个学生去查这个综合分。 学生越多(n 越大),第二种方法的优势越明显。

4.1.2 但是……softmax 怎么办?

等等,上面的推导有个前提:去掉了 softmax。

标准注意力里有个 softmax 函数,它是非线性的。有了 softmax,就不能简单地交换计算顺序了。

那怎么办呢?线性注意力的做法是:把 softmax 去掉,换成一个线性的核函数。

比如,用 φ(x) = elu(x) + 1 这样的函数,让它满足:

1
softmax(QK^T) ≈ φ(Q) φ(K)^T

这样就可以交换顺序了。

⚠️ 注意:线性注意力虽然快,但效果通常不如标准注意力。因为去掉了 softmax,注意力分布的"尖锐度"不够,模型很难把注意力集中在少数几个关键 token 上。

4.1.3 线性注意力的"隐藏技能":RNN 形式

线性注意力还有个很酷的性质:它可以写成循环神经网络(RNN)的形式。

1
2
S_t = S_{t-1} + k_t v_t^T
y_t = q_t^T S_t

什么意思呢?就是你可以一个 token 一个 token 地处理,每一步只需要保存一个状态 S_t。

这有什么好处? - 训练时:用并行形式(Q(K^T V)),可以充分利用 GPU 的并行计算能力 - 推理时:用 RNN 形式,只需要保存一个状态,不需要保存整个 KV Cache

推理速度会非常快,而且内存占用是固定的,不会随序列长度增长!

💡 生活类比: 想象你在读书做笔记: - 标准注意力:每读一句话,都要把前面所有句子翻出来对比一遍。书越厚,越慢。 - 线性注意力(RNN 形式):每读一句话,就更新一下你的"读书笔记"。下次只需要看笔记,不需要翻全书。 笔记的大小是固定的,不管书多厚,都只需要这一本笔记。


4.2 从线性注意力到 Mamba-2

线性注意力很好,但有个问题:效果不够好。

为什么?因为线性注意力对所有输入"一视同仁",没有选择能力。而标准注意力的 softmax 能让模型"聚焦"在重要的 token 上。

那能不能在保持线性复杂度的同时,增加一些选择能力呢?

4.2.1 Mamba-2:加个门控

Mamba-2 的做法很简单:给状态加一个衰减系数 γ_t。

1
2
S_t = γ_t S_{t-1} + k_t v_t^T
y_t = q_t^T S_t + v_t^T D

其中 γ_t = f(x_t),是由输入决定的。

这是什么意思呢?就是说,过去的信息会逐渐"遗忘",而且遗忘的速度由当前输入决定。

  • 如果 γ_t 接近 1:过去的信息保留得多
  • 如果 γ_t 接近 0:过去的信息很快被遗忘

这样模型就可以根据当前输入,动态地决定"要记住多少过去的信息"。

💡 生活类比: 想象你在听讲座,做笔记: - 线性注意力:把所有内容都原封不动记下来,不管重不重要。 - Mamba-2:你有个"遗忘旋钮",听到重要内容时,旋钮调大(多记);听到无关内容时,旋钮调小(快忘)。 旋钮的位置由当前内容决定。

Mamba-2 保持了线性复杂度,同时效果比纯线性注意力好很多。

4.2.2 Gated Delta Net:更进一步的门控

Gated Delta Net(GDN)在 Mamba-2 的基础上又加了一层门控:

1
2
S_t = γ_t (I - β_t k_t k_t^T) S_{t-1} + β_t k_t v_t^T
y_t = q_t^T S_t

这里多了个 β_t,它控制着: 1. 输入的"写入强度":β_t 大,新信息写入多;β_t 小,新信息写入少 2. 状态的"擦除方向":沿着 k_t 的方向擦除旧信息

简单说就是:模型可以选择性地"擦掉"某些方向的旧信息,写入新信息。

这就更灵活了——不仅能控制"记多少",还能控制"记什么、忘什么"。

💡 生活类比: 继续用记笔记的例子: - Mamba-2:你有个总开关,控制整体遗忘速度。 - GDN:你有更精细的控制——可以选择性地擦掉某些科目的旧笔记,同时写入新的内容。 比如学了新的数学知识,就擦掉旧的数学笔记,保留其他科目的。

4.2.3 混合架构:线性 + 标准注意力

这些线性注意力变体虽然快,但效果还是不如标准注意力。

那怎么办呢?混合使用!

比如: - MiniMax M1:7 层线性注意力 + 1 层标准注意力(7:1 混合) - Nemotron 3:3 层 Mamba + 1 层标准注意力(3:1 混合) - Qwen 3.5:3 层 GDN + 1 层标准注意力(3:1 混合)

大部分层用快速的线性注意力,少数几层用标准注意力来保证效果。这样既快又好!

💡 生活类比: 想象一个公司: - 大部分日常工作(70-80%)由普通员工快速处理(线性注意力) - 重要决策(20-30%)由高管会议讨论决定(标准注意力) 这样既有速度,又有质量。


4.3 稀疏注意力:只看该看的

除了线性注意力,还有另一种思路:稀疏注意力。

标准注意力是"每个 token 看所有 token"——这叫稠密注意力。 稀疏注意力是"每个 token 只看一部分 token"——只看该看的。

4.3.1 滑动窗口注意力

最简单的稀疏注意力:每个 token 只看它附近的几个 token(一个窗口内)。

比如窗口大小是 512,那每个 token 只看前后 256 个 token。

计算量从 O(n²) 变成 O(n × w),w 是窗口大小,通常是固定的(比如 512),所以整体是线性的!

💡 生活类比: 想象你在看一部很长的电影: - 标准注意力:每看一个画面,都要回忆整部电影的所有画面。 - 滑动窗口注意力:每看一个画面,只回忆最近几分钟的内容。 大部分时候,你只需要近期的上下文就能理解剧情。

第三讲我们提到过,很多新模型(Cohere、LLaMA 4、Gemma 3/4、OLMo 3)都用了混合注意力——交替使用全注意力和滑动窗口注意力。

4.3.2 DeepSeek Sparse Attention(DSA)

滑动窗口是固定的稀疏模式,但有时候我们需要更灵活的稀疏。

DeepSeek 的 DSA(DeepSeek Sparse Attention)是这样的: - 先用一个轻量级的"索引器"(indexer)找出哪些 token 是重要的 - 然后注意力只在这些重要 token 上计算

这样可以在保持效果的同时,大大减少计算量。

而且 DSA 可以"事后"加上——先在短上下文上预训练一个稠密模型,再加上稀疏注意力来支持长上下文。

💡 生活类比: 想象你在写论文,需要查很多资料: - 标准注意力:把所有相关的书都买回家,一页一页读。 - DSA:先用搜索引擎(索引器)找出最相关的几篇论文,再仔细读这几篇。 又快又好,因为大部分资料其实用不上。


4.4 混合专家模型(MoE):让专家各司其职

上面讲的都是注意力的替代方案。现在我们换个话题:混合专家模型(Mixture of Experts,简称 MoE)。

MoE 不是注意力的变体,而是 FFN 层的变体。

4.4.1 什么是 MoE?

在标准 Transformer 里,每个 FFN 层是一个大的前馈网络,所有 token 都经过同一个 FFN。

MoE 的想法是:把一个大 FFN 换成很多个小 FFN(叫"专家"),然后每个 token 只经过其中少数几个专家。

比如: - 标准模型:1 个大 FFN,所有 token 都走这个 FFN - MoE 模型:8 个小 FFN(专家),每个 token 只选 2 个专家

关键在于:参数量可以很大,但每次计算只激活少数专家,所以 FLOPs 并不多。

💡 生活类比: 想象一家医院: - 标准模型:只有一个"全科医生",什么病都看。这个医生需要懂所有医学知识,压力很大。 - MoE 模型:有很多专科医生(专家)——内科、外科、眼科、皮肤科……每个病人(token)只去看 1-2 个相关的专科医生。 医院总共有很多医生(总参数量大),但每个病人只看少数几个(每次计算的 FLOPs 少)。

4.4.2 为什么 MoE 越来越流行?

MoE 这两年特别火,为什么?

原因一:相同 FLOPs,更多参数,效果更好

研究表明,在相同的计算量下,MoE 模型比稠密模型效果更好。

因为 MoE 可以用更多参数(更多专家),而不需要增加计算量。更多参数通常意味着更强的表达能力。

原因二:训练更快

MoE 模型训练速度更快,因为: - 专家之间可以并行计算 - 每个 token 只经过少数专家,计算量小

原因三:效果有竞争力

现在很多高性能的开源模型都是 MoE: - Mixtral 8x7B - DBRX - DeepSeek V2/V3 - Qwen MoE - Grok

原因四:容易并行

MoE 天然适合多 GPU 并行——每个 GPU 放几个专家,token 在 GPU 之间路由。

💡 为什么之前不流行? MoE 其实是个老想法(2017 年就有了),但之前不流行,因为: 1. 基础设施复杂,不好实现 2. 训练不稳定,容易出问题 3. 效果优势不明显

这几年技术成熟了,效果也验证了,所以就火起来了。


4.5 MoE 的核心问题:路由(Routing)

MoE 最关键的问题是:怎么决定每个 token 去哪个专家?

这就是"路由"(routing)问题。

4.5.1 Top-K 路由

最常见的路由方式是 Top-K 路由:

  1. 用一个简单的线性层(叫"门控"或"路由器")给每个专家打分
  2. 选分数最高的 K 个专家
  3. token 只经过这 K 个专家

K 通常很小: - Switch Transformer:K=1(只选 1 个专家) - GShard、Mixtral、Grok:K=2 - Qwen、DBRX:K=4 - DeepSeek V3:K=8

💡 生活类比: 病人去医院看病: 1. 先去导诊台(路由器),导诊根据症状判断该去哪个科 2. 选最相关的 1-2 个科室(Top-K) 3. 病人只去这几个科室看医生

4.5.2 共享专家

有些 MoE 模型还有"共享专家"——所有 token 都会经过的专家。

比如: - DeepSeek V1:64 个路由专家 + 2 个共享专家 - Qwen 1.5:60 个路由专家 + 4 个共享专家 - DeepSeek V3:256 个路由专家 + 1 个共享专家 - LLaMA 4:128 个路由专家 + 1 个共享专家

共享专家负责处理"通用"的信息,路由专家负责处理"专业"的信息。

💡 生活类比: 医院里,每个病人都要先做基础检查(共享专家),比如量血压、测体温。 然后再去专科医生(路由专家)那里看具体的病。 基础检查是所有人都需要的,专科检查是按需的。

4.5.3 细粒度专家

还有一种设计是"细粒度专家"——把专家做得更小更多。

比如 DeepSeek 的专家是"细粒度"的,每个专家只有标准 FFN 的 1/4 或 1/14 大小。

这样的好处是: - 专家更多,选择更灵活 - 每个专家更小,更容易负载均衡


4.6 MoE 的训练难题

MoE 听起来很好,但训练起来有很多挑战。

4.6.1 问题一:路由不可微

路由决策是"选 Top-K 个专家"——这是个离散的选择,不可微。

什么意思呢?就是说,你没法直接用梯度下降来优化路由器,因为"选哪个专家"这个决策不是连续的。

那怎么办?有几种方法:

方法一:强化学习(RL) 用强化学习来学习路由策略。但 RL 训练不稳定,效果也没好多少,所以现在很少用。

方法二:随机扰动 给路由分数加一点随机噪声,让选择变得"软"一些,这样就有梯度了。 比如 Switch Transformer 用的就是这种方法。

方法三:启发式平衡损失 这是现在最常用的方法。 基本思路:我们希望每个专家被使用的频率大致相等(负载均衡),所以加一个损失函数来鼓励均衡。

最常见的是 Switch Transformer 提出的辅助损失(auxiliary loss): - 如果某个专家被用得太多,就给它惩罚 - 如果某个专家被用得太少,也给它惩罚 - 目标是让所有专家的负载尽量均衡

💡 生活类比: 医院里,如果所有病人都去看同一个医生,那个医生会累死,其他医生又没事干。 所以导诊台(路由器)要尽量把病人均匀分配给各个医生。 辅助损失就像一个"考核指标"——如果分配不均,就扣奖金。

4.6.2 问题二:负载均衡

MoE 训练中最头疼的问题就是负载均衡。

理想情况:每个专家处理的 token 数量差不多。 实际情况:经常出现"热门专家"和"冷门专家"——有的专家忙死,有的专家闲死。

为什么这是个问题? - 系统效率低:闲的专家在浪费算力,忙的专家成了瓶颈 - 训练不稳定:冷门专家见的数据太少,学不好

怎么解决? - 辅助损失(上面说的) - 每个设备的负载均衡(DeepSeek 的 per-device balancing) - 每个专家的偏置项(DeepSeek V3 的 aux-loss-free balancing)

DeepSeek V3 甚至提出了"无辅助损失"的方法——给每个专家加一个可学习的偏置,用在线学习的方式来调整负载。


4.7 MoE 的其他问题

4.7.1 稳定性问题

MoE 训练比稠密模型更容易不稳定。

为什么?因为路由器的输出如果很大,softmax 之后会很尖锐,导致梯度消失或爆炸。

解决方法: - 路由器用 FP32 精度(其他部分用 BF16) - 给路由器加 Z-Loss(第三讲学过的)

Z-Loss 就是让路由器的输出 logits 不要太大,保持在合理范围内。

4.7.2 微调问题

稀疏 MoE 模型在小数据集上微调时,容易过拟合。

为什么?因为 MoE 参数量很大,小数据集不够训。

解决方法: - 微调时只用非 MoE 的部分(比如只微调注意力层) - 用更多的微调数据(DeepSeek 用了 1.4M SFT 数据)

4.7.3 随机性问题

MoE 模型还有个有趣的问题:结果可能有随机性。

为什么?因为 token 丢弃(token dropping)是在 batch 级别进行的。

什么意思呢?如果一个专家太忙了,超出容量的 token 会被丢弃。但哪些 token 被丢弃,取决于同一个 batch 里其他 token 的选择。

也就是说,你的输入的处理结果,可能会被同一个 batch 里别人的输入影响!

这听起来有点奇怪,但确实是 MoE 的一个特点。


4.8 Upcycling:把稠密模型变成 MoE

还有个有趣的技术叫 Upcycling(升级改造)。

什么意思呢?就是拿一个已经训练好的稠密模型,把它改造成 MoE 模型。

怎么做? 1. 把原来的 FFN 层复制多份,变成多个专家 2. 加一个路由器 3. 继续训练一段时间

这样做的好处: - 不需要从零开始训练 MoE,节省时间 - 可以利用已有的稠密模型的知识

成功案例: - MiniCPM MoE:从 MiniCPM 稠密模型升级而来 - Qwen MoE:从 Qwen 1.8B 稠密模型升级而来

💡 生活类比: 想象你有一家小餐馆(稠密模型),生意很好,想扩大规模。 - 方法一:从零开始建一家大餐厅(从零训练 MoE) - 方法二:把原来的厨师培养成主厨,再招几个专科厨师(川菜、粤菜、日料……),把原来的厨房改造成"美食广场"(Upcycling) 第二种方法更快,因为基础已经有了。


4.9 案例:DeepSeek MoE 系列

让我们用 DeepSeek 的 MoE 系列作为案例,看看 MoE 是怎么演进的。

DeepSeek V1

  • 总参数 16B,激活参数 2.8B
  • 64 个细粒度专家(每个是标准的 1/4)+ 2 个共享专家
  • 每个 token 选 6 个专家
  • 标准的辅助损失平衡(专家级 + 设备级)

DeepSeek V2

  • 总参数 236B,激活参数 21B
  • 160 个细粒度专家(每个是标准的 1/10)+ 2 个共享专家
  • 每个 token 选 6 个专家
  • 新东西:通信平衡损失、Top-M 设备路由

DeepSeek V3

  • 总参数 671B,激活参数 37B
  • 256 个细粒度专家 + 1 个共享专家
  • 每个 token 选 8 个专家
  • 新东西:无辅助损失平衡、序列级辅助损失

可以看到,DeepSeek 的 MoE 越来越大,专家越来越多,激活参数也越来越大。


4.10 Bonus:MLA 和 MTP

DeepSeek V3 还有两个重要的技术,虽然不是 MoE 本身,但经常和 MoE 一起出现。

4.10.1 MLA(Multi-head Latent Attention)

MLA 是一种压缩 KV Cache 的方法。

核心思想:把 K 和 V 表示成一个更低维度的"潜在"激活的函数。

什么意思呢?就是说,我们不保存完整的 K 和 V,只保存一个很小的"压缩版"(叫 c_t)。需要的时候,再从这个压缩版恢复出 K 和 V。

好处:KV Cache 可以小很多!

比如,原来 KV Cache 需要存 512 维的 K 和 512 维的 V,现在只需要存 64 维的 c_t,节省了很多内存。

💡 生活类比: 想象你要保存很多照片: - 标准注意力:保存原始高清照片(KV Cache 大) - MLA:保存压缩后的缩略图(c_t),需要看的时候再放大恢复。 压缩版占用空间小很多,但画质(效果)损失不大。

不过 MLA 有个问题:和 RoPE(旋转位置编码)不太兼容。因为 RoPE 需要旋转 K,而压缩后的 K 不好旋转。

解决方法:保留少数几个非压缩的维度,专门用来做 RoPE。

4.10.2 MTP(Multi-Token Prediction)

MTP 是一种加速推理的方法。

核心思想:让模型一次预测多个 token,而不是一个。

具体做法:在主模型旁边加几个小的"轻量级模型",每个负责预测接下来的一个 token。

比如: - 主模型预测第 t 个 token - 小模型 1 预测第 t+1 个 token - 小模型 2 预测第 t+2 个 token

这样推理的时候,可以一次生成多个 token,速度更快。

💡 生活类比: 想象你在打字: - 标准方式:打一个字,想一下,打下一个字。 - MTP:打一个字,同时预测接下来几个字,一起打出来。 速度更快,因为减少了"思考"的次数。

MTP 和 EAGLE(另一种投机解码方法)思路类似,都是用小模型来加速大模型的推理。


4.11 本讲小结

这一讲我们学了两种让大模型更高效的思路:

1. 注意力替代方案 - 线性注意力:O(n) 复杂度,可写成 RNN 形式,推理快 - Mamba-2:加门控的线性注意力,效果更好 - GDN:更精细的门控,选择性地读写状态 - 混合架构:大部分层用线性注意力,少数层用标准注意力 - 稀疏注意力:只看重要的 token,减少计算量

2. 混合专家模型(MoE) - 核心思想:多个专家,每个 token 只选少数几个 - 优势:相同 FLOPs,更多参数,效果更好 - 关键问题:路由(怎么选专家)、负载均衡、训练稳定性 - 进阶技术:共享专家、细粒度专家、Upcycling - 相关技术:MLA(压缩 KV Cache)、MTP(多 token 预测)


4.12 小白常见问题 Q&A

Q1:线性注意力和标准注意力,到底差多少?

A:效果上确实有差距,但差距在缩小。纯线性注意力效果一般,但加上门控(Mamba、GDN)之后好了很多。再加上混合架构(大部分线性 + 少数标准注意力),效果可以接近纯标准注意力,同时速度快很多。

Q2:MoE 这么好,为什么还有稠密模型?

A:MoE 不是银弹,它有缺点: 1. 基础设施复杂,实现和调试都难 2. 训练不稳定,需要很多技巧 3. 推理时需要加载所有专家,内存占用大 4. 小模型上优势不明显

小模型(比如 7B 以下)通常还是稠密的更划算。大模型(几十 B 以上)用 MoE 更有优势。

Q3:MoE 的"激活参数"是什么意思?

A:MoE 有两个参数数字: - 总参数:所有专家的参数加起来 - 激活参数:每个 token 实际经过的专家的参数

比如 DeepSeek V3 总参数 671B,但每个 token 只激活 37B 的参数。

为什么要区分?因为计算量(FLOPs)只和激活参数有关,和总参数无关。总参数大不代表计算慢。

Q4:为什么 MoE 推理慢?

A:MoE 推理有几个挑战: 1. 需要加载所有专家到显存,内存占用大 2. token 路由需要通信(如果专家在不同 GPU 上) 3. 专家负载不均衡,有的 GPU 忙有的闲 4. 批处理(batching)更难,因为每个 token 走的路径不一样

所以 MoE 训练有优势,但推理不一定有优势。


4.13 学习路线图

必须掌握: - 注意力的 O(n²) 问题是什么 - 线性注意力的核心思想(交换计算顺序) - MoE 的基本概念(多个专家,选 Top-K) - MoE 的优势和挑战 - 路由和负载均衡是什么

了解即可: - Mamba-2 和 GDN 的具体公式 - 各种 MoE 变体(共享专家、细粒度专家) - Upcycling 的概念 - MLA 和 MTP 的基本思想

配合作业: - 作业 2(Systems)里有 Triton kernel 编程,可以配合 L5-L6 一起学 - 作业 5(Alignment)里可能会用到 MoE 模型,可以了解一下


第五讲:GPU 系统与性能优化

5.0 前置知识:GPU 到底是什么?

在第二讲我们简单提过 GPU,说它是"做矩阵乘法很快的芯片"。这一讲我们要深入了解:GPU 到底是怎么工作的?为什么它做深度学习这么快?

先回顾一下 CPU 和 GPU 的区别:

CPU GPU
设计目标 低延迟(快速完成单个任务) 高吞吐量(同时处理很多任务)
核心数量 几个到几十个 几千到上万个
缓存 大(为了减少延迟) 小(为了放更多计算单元)
擅长 复杂的串行任务 简单的并行任务

💡 生活类比: - CPU 像一个经验丰富的老医生,什么病都能看,但一次只能看一个病人。 - GPU 像一个有很多实习医生的医院,每个医生只会做简单的检查,但可以同时看几百个病人。

深度学习的任务(比如矩阵乘法)就像是"给几百个病人量血压"——简单但量大,特别适合 GPU。


5.1 GPU 的内部结构

5.1.1 计算单元:SM 和 SP

GPU 不是一个大芯片,而是由很多个流式多处理器(Streaming Multiprocessor,简称 SM)组成的。

每个 SM 里面又有很多个流式处理器(Streaming Processor,简称 SP),也叫 CUDA 核心。

结构大概是这样的:

1
2
3
4
5
6
7
8
9
GPU
├── SM 0
│ ├── SP 0
│ ├── SP 1
│ ├── ...
│ └── SP 127
├── SM 1
├── ...
└── SM 107

比如 A100 GPU 有 108 个 SM,每个 SM 有很多个 SP。

每个 SM 可以独立执行任务,SP 是真正做计算的地方。

💡 生活类比: GPU 像一个工厂: - 整个工厂(GPU)有很多车间(SM) - 每个车间(SM)有很多工人(SP) - 每个工人(SP)做简单的计算 - 车间之间相对独立,可以同时干不同的活

5.1.2 内存层次结构:离得越近越快

GPU 的内存是分层的,离计算单元越近,速度越快,但容量越小:

内存类型 位置 速度 容量(A100)
寄存器(Register) 每个线程自己的 最快 每个 SM 256 KB
共享内存 / L1 缓存 SM 内部 很快 每个 SM 192 KB
L2 缓存 GPU 芯片上 较快 40 MB
全局内存(HBM) GPU 旁边的显存芯片 较慢 80 GB

速度差距有多大呢?大概是这样的: - 寄存器带宽:~116 TB/s - 共享内存带宽:~19 TB/s - L2 缓存带宽:~5-8 TB/s - HBM 带宽:~2 TB/s

寄存器比 HBM 快 50 多倍!

💡 生活类比: 内存层次就像你的书桌: - 寄存器 = 你手里拿着的笔,随时能用,但只能拿几支 - 共享内存/L1 = 桌面上的东西,伸手就能拿到,能放不少 - L2 缓存 = 书架上的书,站起来就能拿,能放更多 - HBM(全局内存)= 楼下图书馆的书,要下楼去借,能放非常多,但慢

为了效率,你应该尽量多用桌面上的东西,少跑图书馆。


5.2 GPU 的执行模型

5.2.1 线程、块、网格

GPU 编程里有三个重要概念:

  1. 线程(Thread):最小的执行单元,每个线程处理一小部分数据
  2. 线程块(Thread Block):一组线程,共享同一个 SM 和共享内存
  3. 网格(Grid):所有线程块的集合

结构是这样的:

1
2
3
4
5
6
7
8
9
Grid(网格)
├── Block 0(线程块)
│ ├── Thread 0
│ ├── Thread 1
│ ├── ...
│ └── Thread 127
├── Block 1
├── ...
└── Block 999

为什么要分块?因为线程之间需要通信(比如计算总和),而通信需要共享内存。同一个块里的线程可以通过共享内存快速通信,不同块之间不能直接通信。

💡 生活类比: 想象一个大工厂要组装 10000 个产品: - 每个产品由一个工人(线程)组装 - 工人们分成小组(线程块),每个小组在一个车间(SM)里工作 - 同一个小组的工人可以互相递工具(共享内存) - 不同小组之间不能直接交流 - 所有小组加起来就是整个工厂的生产任务(网格)

5.2.2 Warp:32 个线程一组

还有个重要概念叫 Warp( warp)。

在 GPU 里,线程不是一个一个执行的,而是 32 个线程一组,这一组就叫一个 Warp。

同一个 Warp 里的所有线程,必须执行相同的指令——这叫 SIMT(Single Instruction, Multiple Threads,单指令多线程)。

什么意思呢?就是说,32 个线程步调一致,同时做同样的操作,只是处理的数据不同。

💡 生活类比: 想象一个合唱团: - 32 个歌手(线程)组成一个声部(Warp) - 所有人同时唱同一个音符(同一条指令) - 但每个人的音色略有不同(处理的数据不同) - 指挥(SM)一次指挥一个声部

5.2.3 控制分支:为什么 if-else 可能很慢?

SIMT 模型有个问题:如果同一个 Warp 里的线程需要执行不同的指令怎么办?

比如有个 if-else 语句: - 线程 0-15 满足条件,走 if 分支 - 线程 16-31 不满足条件,走 else 分支

这时候 GPU 会怎么做?答案是:先一起执行 if 分支(else 分支的线程等着),再一起执行 else 分支(if 分支的线程等着)。

这就叫 控制分支(Control Divergence)。

如果分支很多,Warp 里的线程各走各的,那效率就会非常低——因为大部分时间大家都在等别人。

💡 生活类比: 合唱团唱歌: - 如果所有人唱同一首歌,效率很高 - 如果一半人唱《小星星》,另一半人唱《两只老虎》,那就得先唱一半的歌,再唱另一半的歌,时间翻倍 - 如果每个人唱不同的歌,那就彻底乱了,效率极低

所以 GPU 编程要尽量避免分支,或者让同一个 Warp 里的线程走相同的分支。


5.3 Roofline 模型:怎么判断程序快不快?

第二讲学过算术强度和 Roofline 模型,这里再复习一下,因为这是理解 GPU 性能的核心。

5.3.1 两种瓶颈

一个 GPU 程序的性能,通常受两种因素限制:

  1. 计算受限(Compute Bound):计算单元不够用,内存带宽还有富余
  2. 内存受限(Memory Bound):内存带宽不够用,计算单元在闲着等数据

怎么判断是哪种?看算术强度:

1
算术强度 = 计算量(FLOPs) / 内存访问量(字节)
  • 算术强度高 → 计算受限(计算是瓶颈)
  • 算术强度低 → 内存受限(内存是瓶颈)

5.3.2 Roofline 图

把这个关系画成图,就是 Roofline 模型:

1
2
3
4
5
6
7
8
9
10
11
性能 (FLOP/s)
↑
| ┌───────────────────── 计算峰值(屋顶)
| /
| /
| / ← 内存带宽线(斜率 = 带宽)
| /
| /
| /
| /
+--------------------------→ 算术强度 (FLOPs/字节)
  • 左边斜线部分:内存受限,算术强度越低,性能越低
  • 右边平顶部分:计算受限,性能达到峰值,不再增长

你的程序的算术强度落在哪个区域,就决定了它的瓶颈是什么。

💡 生活类比: 想象一个工厂: - 工人(计算单元)的生产速度是有限的(计算峰值) - 原材料运输(内存带宽)的速度也是有限的 - 如果产品需要很多原材料(算术强度低),那运输跟不上,工人闲着等(内存受限) - 如果产品需要很少原材料(算术强度高),那工人全力生产,运输不是问题(计算受限)

优化的目标:要么提高算术强度(减少内存访问),要么提高计算效率。


5.4 GPU 性能优化的六大技巧

了解了 GPU 的工作原理,我们来看看怎么让程序跑得更快。

5.4.1 技巧一:低精度计算

最简单的优化:用更少的位数来表示数字。

比如: - FP32(32 位浮点数):4 字节 - FP16/BF16(16 位浮点数):2 字节 - FP8(8 位浮点数):1 字节 - FP4(4 位浮点数):0.5 字节

为什么低精度更快? 1. 内存访问少了:同样的数据,位数越少,字节数越少,读得越快 2. 计算更快:低精度的计算单元(比如 Tensor Core)比普通计算单元快很多 3. 显存占用少了:同样的显存能放更多参数

这就是为什么混合精度训练(第二讲学过)这么重要——用 BF16 做大部分计算,既快又够用。

💡 生活类比: 想象你要运一批货物: - 高精度 = 每个货物用大木箱装,占地方,运得慢 - 低精度 = 每个货物用小纸盒装,占地方小,运得快

如果货物不需要那么精密的保护,用小纸盒就够了,效率高很多。

前沿:FP8 和 FP4

现在最新的 GPU(比如 B200)已经支持 FP8 甚至 FP4 了。

FP8 有两种格式: - E4M3:4 位指数,3 位尾数——动态范围大,精度稍低 - E5M2:5 位指数,2 位尾数——动态范围更大,精度更低

还有更激进的 MXFP4——4 位浮点数,每 16 个元素共享一个缩放因子。

精度越低,速度越快,但效果可能会下降。需要找到平衡点。


5.4.2 技巧二:算子融合(Operator Fusion)

第二个重要技巧:把多个操作合并成一个。

为什么要融合?因为每个操作都要读数据、写数据,来回读写内存很费时间。如果把多个操作合并成一个,只需要读一次、写一次,节省很多内存访问。

举个例子:计算 sin²(x) + cos²(x)。

朴素做法: 1. 读 x,计算 sin(x),写回内存 2. 读 sin(x),计算平方,写回内存 3. 读 x,计算 cos(x),写回内存 4. 读 cos(x),计算平方,写回内存 5. 读 sin²(x) 和 cos²(x),相加,写回内存

总共 5 次读、5 次写!

融合做法: 1. 读 x 2. 一次计算 sin(x)、cos(x)、平方、相加 3. 写结果

只需要 1 次读、1 次写!

💡 生活类比: 想象你在厨房做菜: - 不融合:切菜 → 装盘 → 端到灶台 → 炒菜 → 装盘 → 端到餐桌 → 加调料 → 装盘…… 来回跑,效率低 - 融合:在灶台边,切完直接炒,炒完直接加调料,一气呵成,少跑很多路

减少来回搬运(内存访问),效率就高了。

PyTorch 2.0 的 torch.compile() 就能自动做很多算子融合,把多个小操作合并成一个大的 CUDA kernel。


5.4.3 技巧三:重计算(Recomputation)

第三个技巧有点反直觉:扔掉中间结果,需要的时候重新算。

为什么这样反而更快?因为保存中间结果需要写内存,读取中间结果需要读内存,如果重计算的代价比内存访问的代价小,那重算反而更快。

最典型的例子就是激活检查点(Activation Checkpointing)——第二讲学过。

训练时,前向传播会产生很多激活值,反向传播需要这些激活值来算梯度。如果全部保存,显存占用很大。

激活检查点的做法是:只保存部分层的激活值,其他层的激活值在反向传播时重新计算。

这样用计算换内存——多花一些计算时间,但节省了大量显存。

💡 生活类比: 想象你在做数学题: - 不重计算:每一步的中间结果都写在纸上,最后检查的时候直接看。纸(显存)用得多。 - 重计算:只写关键步骤的结果,中间步骤不写。检查的时候,从关键步骤重新算中间结果。纸用得少,但要多算一遍。

如果纸很贵(显存紧张),那重算是划算的。


5.4.4 技巧四:内存合并(Memory Coalescing)

第四个技巧:让同一个 Warp 的线程访问连续的内存地址。

为什么?因为 GPU 的内存是按"突发模式(burst mode)"读的——一次读一大块(比如 128 字节)。

如果同一个 Warp 的 32 个线程访问的地址是连续的,那只需要一次内存读取就能拿到所有数据。

如果地址是分散的,那就需要很多次读取,速度就慢了。

💡 生活类比: 想象你去图书馆借书: - 合并访问:你借的 32 本书都在同一排书架上,管理员一次就能拿过来。 - 不合并访问:你借的 32 本书分散在不同的楼层不同的书架,管理员要跑很多趟。

连续访问效率高,分散访问效率低。

矩阵乘法里,行优先存储的矩阵,按行读是合并的,按列读是不合并的。这就是为什么矩阵乘法的实现要仔细考虑内存访问模式。


5.4.5 技巧五:分块(Tiling)

第五个技巧,也是最重要的一个:分块。

分块的思想是:把大矩阵分成小块(tile),让小块能放进共享内存里,这样就可以重复利用数据,减少全局内存访问。

以矩阵乘法 C = A × B 为例:

朴素做法: - 每个元素 C[i,j] 需要 A 的第 i 行和 B 的第 j 列 - 每次都从全局内存读 A 和 B - 同一个行/列被重复读很多次,浪费带宽

分块做法: 1. 把 C 分成很多小块(比如 64×64) 2. 每个小块由一个线程块负责 3. 计算时,把 A 和 B 的对应小块加载到共享内存 4. 在共享内存里做计算,结果写回全局内存

这样 A 和 B 的每个元素只需要从全局内存读一次(或者很少几次),大大减少了全局内存访问。

💡 生活类比: 想象你要组装很多家具: - 不分块:每次组装一个零件,都去仓库(全局内存)拿原材料。来回跑,效率低。 - 分块:先把一批需要的原材料都搬到车间(共享内存),然后在车间里组装。少跑很多趟仓库。

数据复用是关键——能在快速内存里搞定的,就不要去慢速内存。

分块是 GPU 性能优化中最重要的技巧之一。几乎所有高性能的矩阵乘法实现都用了分块。


5.4.6 技巧六:提高占用率(Occupancy)

最后一个技巧:提高 SM 的占用率。

什么是占用率?就是一个 SM 上同时能跑多少个 Warp。

占用率受什么限制? - 寄存器:每个线程用的寄存器越多,能同时跑的线程越少 - 共享内存:每个块用的共享内存越多,能同时跑的块越少 - Warp 数量上限:每个 SM 最多能跑多少个 Warp(比如 64 个)

占用率低有什么问题? - 当一个 Warp 在等内存的时候,SM 可以切换到另一个 Warp 继续算 - 如果 Warp 太少,SM 就会经常闲着等内存

所以占用率高一点好,但也不是越高越好——如果每个线程的工作量大,低占用率也能跑得很快。

💡 生活类比: 想象一个车间(SM): - 工人(Warp)越多,当一组工人在等材料(内存访问)时,可以换另一组工人继续干。 - 但如果工人太多,每个人的工具(寄存器)不够用,反而效率低。 - 需要找到平衡点。


5.5 案例分析:FlashAttention

学了这么多优化技巧,我们来看一个经典案例:FlashAttention。

FlashAttention 是 2022 年提出的一种快速注意力实现,比标准注意力快很多,现在几乎所有大模型都在用。

5.5.1 标准注意力的问题

标准注意力的计算过程是: 1. 计算 S = QK^T / √d_k(缩放后的注意力分数矩阵) 2. 计算 P = softmax(S) 3. 计算 O = PV

问题在哪?S 和 P 都是 n×n 的大矩阵,需要读写全局内存,很慢。

特别是第二步的 softmax,需要先读 S,计算后写 P,然后第三步又要读 P——来回读写很费时间。

5.5.2 FlashAttention 的思路

FlashAttention 的核心思想就是我们刚学的两个技巧:分块 + 算子融合。

具体怎么做?

  1. 分块:把 Q、K、V 都分成小块
  2. 逐块计算:每次加载一小块 Q 和一小块 K V 到共享内存
  3. 在线 softmax:逐块计算 softmax,不需要保存完整的注意力矩阵
  4. 融合:整个过程在一个 kernel 里完成,减少内存访问

关键是在线 softmax(online softmax)——怎么在不保存完整矩阵的情况下计算 softmax?

softmax 的公式是:

1
softmax(x_i) = exp(x_i) / Σ exp(x_j)

问题是,你需要知道所有 x_j 的和才能归一化,但如果数据是分块来的,你不知道后面的块有多大。

解决方法:增量更新最大值 + telescoping sum(望远镜求和)。

简单说就是: 1. 每来一个新块,更新当前看到的最大值 2. 因为最大值变了,之前算的指数都要缩放 3. 用一个巧妙的方式来累积,保证最后结果是对的

这样就可以逐块计算 softmax,不需要保存完整的 n×n 矩阵。

💡 生活类比: 想象你要算全班同学的平均身高: - 标准方法:先把所有人的身高都记下来(保存完整矩阵),再加起来除以人数。 - FlashAttention 方法:同学一个个进来(分块),你一边看一边更新"目前最高身高"和"目前的总和",最后算平均。 - 不需要记住所有人的身高,只需要记住几个统计量。

FlashAttention 把注意力计算的内存访问从 O(n²) 降到了 O(n),同时因为融合了多个操作,速度快了很多倍。


5.6 本讲小结

这一讲我们深入了解了 GPU 的工作原理和性能优化技巧:

GPU 基础 - GPU 由很多 SM 组成,每个 SM 有很多 SP - 内存是分层的,离得越近越快 - 执行模型:Thread → Block → Grid,32 个线程组成一个 Warp - SIMT 模型:同一个 Warp 的线程执行相同指令

性能优化六大技巧 1. 低精度计算:减少数据量,提高计算速度 2. 算子融合:合并多个操作,减少内存访问 3. 重计算:用计算换内存,减少显存占用 4. 内存合并:让线程访问连续地址,提高内存效率 5. 分块:把数据分成小块,利用共享内存复用数据 6. 提高占用率:让 SM 保持忙碌

案例:FlashAttention - 用分块 + 在线 softmax + 融合,大大加速注意力计算 - 内存访问从 O(n²) 降到 O(n)


5.7 小白常见问题 Q&A

Q1:GPU 和 TPU 有什么区别?

A:GPU 和 TPU 都是深度学习加速器,核心思想类似——大量并行计算单元 + 快速内存。但也有区别: - GPU 更通用,什么都能算,编程更灵活 - TPU 更专门,主要做矩阵乘法,在特定任务上效率更高 - 网络连接方式不同:TPU 是环形网格(toroidal mesh),GPU 是全连接 - 编程模型不同:GPU 用 CUDA/Triton,TPU 用 JAX/XLA

简单说,GPU 是"瑞士军刀",TPU 是"专业手术刀"。

Q2:为什么矩阵乘法这么重要?

A:因为深度学习里大部分计算都是矩阵乘法。Transformer 里的注意力、FFN,全都是矩阵乘法。

而且矩阵乘法的算术强度很高(大矩阵乘法是计算受限的),特别适合 GPU 的 Tensor Core 来加速。

所以优化深度学习性能,很大程度上就是优化矩阵乘法。

Q3:Triton 和 CUDA 是什么关系?

A:CUDA 是 NVIDIA 推出的 GPU 编程语言,很底层,控制力强,但写起来复杂。

Triton 是 OpenAI 推出的高级 GPU 编程语言,比 CUDA 简单,更容易写高性能的 kernel。

Triton 最终也会编译成 PTX(GPU 的汇编语言),和 CUDA 本质上是一样的,只是抽象层次更高。

下一讲我们会详细学 Triton。

Q4:FlashAttention 为什么这么快?

A:主要有三个原因: 1. 减少内存访问:用分块的方式,不需要保存完整的 n×n 注意力矩阵,内存访问从 O(n²) 降到 O(n) 2. 算子融合:把 QK 相乘、softmax、乘 V 都融合在一个 kernel 里,减少了中间结果的读写 3. 利用共享内存:小块数据放在共享内存里,重复利用,减少全局内存访问

这三个都是我们学过的优化技巧的组合应用。


5.8 学习路线图

必须掌握: - GPU 的基本结构(SM、SP、内存层次) - Warp 和 SIMT 的概念 - 控制分支为什么会影响性能 - Roofline 模型和算术强度 - 六大优化技巧的基本思想 - FlashAttention 的核心思路

了解即可: - 具体的 CUDA/Triton 编程细节 - 各种低精度格式的区别(E4M3 vs E5M2 等) - 在线 softmax 的具体推导

配合作业: - 作业 2(Systems)里会让你写 Triton kernel,可以配合 L6 一起学 - 作业 2 里也有 Flash Attention 的实现,可以结合这一讲的内容理解


第六讲:Triton 编程与性能调优

6.0 前置知识:为什么要学 Triton?

上一讲我们学了 GPU 的工作原理和优化技巧。这一讲我们要动手写 GPU 程序。

你可能会问:PyTorch 不是已经帮我们写好了 GPU 代码吗?为什么还要自己写?

答案是:有时候 PyTorch 自带的算子不够快,或者没有你需要的算子。

比如: - 你想把多个操作融合成一个 kernel,减少内存访问 - 你想实现一个新的算法,PyTorch 没有现成的 - 你想针对特定的硬件做深度优化

这时候就需要自己写 GPU kernel 了。

写 GPU kernel 有两种选择: 1. CUDA:NVIDIA 的官方语言,很底层,控制力强,但写起来复杂 2. Triton:OpenAI 推出的高级语言,比 CUDA 简单,更容易写高性能代码

这门课用的是 Triton,因为它更适合初学者,也足够强大。

💡 生活类比: - PyTorch 自带算子 = 超市里的现成菜,打开就能吃,但不一定合你口味 - 自己写 Triton kernel = 自己下厨,可以按自己的口味做,还能把几道菜融合成一道 - CUDA = 从种菜开始自己做,完全可控,但太麻烦

Triton 是一个很好的平衡点——足够灵活,又不至于太复杂。


6.1 GPU 硬件回顾:三代 GPU 对比

先快速回顾一下 GPU 硬件,看看不同代的 GPU 有什么区别。

指标 A100 H100 B200
SM 数量 108 132 148
寄存器(每个 SM) 256 KB 256 KB 256 KB
L1 缓存 + 共享内存 192 KB 256 KB 256 KB
L2 缓存 40 MB 50 MB 96-126 MB
HBM 大小 80 GB 80 GB 192 GB
寄存器带宽 ~116 TB/s ~401 TB/s ~447 TB/s
共享内存带宽 ~19 TB/s ~33 TB/s ~19 TB/s
L2 缓存带宽 ~5-8 TB/s ~12 TB/s ~9 TB/s
HBM 带宽 ~2 TB/s 3.35 TB/s ~8 TB/s

可以看到: - 每一代 GPU 的算力和内存带宽都在提升 - 内存层次结构(寄存器 → 共享内存 → L2 → HBM)一直没变 - 离计算单元越近,速度越快——这个规律也没变

💡 为什么要了解这些数字? 因为优化 GPU 程序的核心就是减少慢速内存的访问,增加快速内存的访问。 知道每层内存的速度和容量,你才能判断你的程序的瓶颈在哪,该怎么优化。


6.2 编程模型:Thread、Block、Grid

上一讲学过,GPU 的编程模型是三层的:

  1. Thread(线程):最小的执行单元,每个线程处理一小部分数据
  2. Thread Block(线程块):一组线程,共享同一个 SM 和共享内存
  3. Grid(网格):所有线程块的集合
1
2
3
4
5
6
7
8
9
Grid
├── Block 0
│ ├── Thread 0
│ ├── Thread 1
│ ├── ...
│ └── Thread 127
├── Block 1
├── ...
└── Block N

为什么要分块?

因为线程之间需要通信(比如计算总和),而通信需要共享内存。同一个块里的线程可以通过共享内存快速通信,不同块之间不能直接通信。

Triton 的特点

Triton 和 CUDA 有个重要区别: - CUDA:你指定每个线程做什么 - Triton:你指定每个线程块做什么

Triton 的抽象层次更高——你不需要管单个线程怎么工作,只需要管整个块怎么工作。Triton 编译器会自动把块的工作分配给线程。

这就是为什么 Triton 更容易写——你不用操心太多底层细节。

💡 生活类比: - CUDA:你给每个工人(线程)分配具体任务 - Triton:你给每个小组(线程块)分配整体任务,组长(编译器)负责给工人分工

Triton 更省心,因为很多细节编译器帮你搞定了。


6.3 硬件细节:影响性能的关键因素

了解了编程模型,我们来看看几个影响性能的关键硬件细节。

6.3.1 Warp 和控制分支

上一讲学过,32 个线程组成一个 Warp,同一个 Warp 的线程执行相同的指令(SIMT)。

如果同一个 Warp 里的线程走不同的分支(if-else),就会有控制分支(Control Divergence),效率下降。

优化建议: - 尽量避免分支 - 如果必须有分支,尽量让同一个 Warp 里的线程走相同的分支

6.3.2 Warp 占用率(Occupancy)

每个 SM 能同时跑多少个 Warp?这叫占用率(Occupancy)。

占用率受什么限制? - 寄存器:每个线程用的寄存器越多,能同时跑的线程越少 - 共享内存:每个块用的共享内存越多,能同时跑的块越少 - Warp 数量上限:每个 SM 最多能跑 64 个 Warp(A100)

举个例子: - 每个块 128 个线程 = 4 个 Warp - 每个线程用 160 个寄存器 - 每个 SM 有 65536 个寄存器 - 每个块用 128 × 160 = 20480 个寄存器 - 每个 SM 能跑 65536 / 20480 ≈ 3 个块 = 12 个 Warp - 占用率 = 12 / 64 = 18.75%

这个占用率很低,意味着 SM 经常会闲着等内存。

怎么提高占用率? - 减少每个线程的寄存器使用量 - 减小块的大小 - 减少共享内存的使用

但也不是占用率越高越好——如果每个线程的工作量很大,低占用率也能跑得很快。

💡 生活类比: 想象一个车间(SM): - 工人(Warp)越多,当一组工人在等材料时,可以换另一组继续干 - 但如果每个工人需要很多工具(寄存器),车间里放不下那么多工人 - 需要找到平衡点——工人不能太少(闲着等),也不能太多(工具不够)

6.3.3 Bank 冲突(Bank Conflicts)

共享内存虽然快,但有个坑——Bank 冲突。

共享内存被分成 32 个 bank(存储体),每个 bank 4 字节宽。每个周期,每个 bank 只能被一个线程访问。

如果同一个 Warp 的 32 个线程访问同一个 bank 的不同地址,那就会冲突——访问会串行化,速度变慢。

最坏情况:32 个线程都访问同一个 bank,变成 32 次串行访问,速度慢 32 倍!

怎么避免 Bank 冲突? - 让线程访问的地址分布在不同的 bank - 用 padding(填充)错开地址 - 用 swizzling(交错排列)重排数据

矩阵乘法里,读 A 的行和读 B 的列,很容易产生 Bank 冲突,需要特别注意。

💡 生活类比: 想象一个银行(共享内存)有 32 个窗口(bank): - 如果 32 个人(线程)分别去 32 个窗口办业务,大家同时办,很快 - 如果 32 个人都去同一个窗口,就得排队,很慢 - Bank 冲突就是大家挤在同一个窗口

6.3.4 内存合并(Memory Coalescing)

全局内存(HBM)的访问也讲究——同一个 Warp 的线程访问的地址如果是连续的,就能合并成一次内存事务,速度快。

如果地址是分散的,就需要多次事务,速度慢。

优化建议: - 尽量让同一个 Warp 的线程访问连续的内存地址 - 矩阵按行存储时,按行访问是合并的,按列访问是不合并的

6.3.5 块占用率(Block Occupancy)

最后一个概念:块占用率。

GPU 有很多 SM(比如 A100 有 108 个),你的网格有 N 个块。

如果 N 正好是 SM 数量的整数倍,那每个 SM 跑的块数一样,大家同时完成,效率高。

如果 N 不是整数倍,那最后一波只有少数几个 SM 在跑,其他 SM 闲着——这叫波量化(Wave Quantization)问题。

比如: - 108 个 SM,160 个块 - 第一波:108 个块同时跑 - 第二波:剩下 52 个块,只有 52 个 SM 在工作,另外 56 个 SM 闲着

这就浪费了。

优化建议: - 尽量让块的数量是 SM 数量的整数倍 - 或者让块足够多,最后一波的影响可以忽略


6.4 Benchmarking 和 Profiling

写 GPU 程序,第一步不是优化,而是测量——先搞清楚哪里慢,再针对性优化。

有两种测量工具:

6.4.1 Benchmarking(基准测试)

Benchmarking 测量的是端到端的时间——整个操作花了多少毫秒。

它告诉你"快不快",但不告诉你"为什么快/慢"。

Benchmarking 可以用来: - 比较不同实现的速度 - 看性能随数据规模怎么变化

PyTorch 里可以用 torch.utils.benchmark,也可以自己写。

怎么正确地 benchmark? 1. Warmup(预热):先跑几次,让编译、缓存等准备工作完成 2. 多次运行取平均:单次运行有波动,多跑几次取平均 3. 用 CUDA Event 计时:不要用 CPU 的时间,要用 GPU 的事件来计时,因为 CPU 和 GPU 是异步的

⚠️ 注意:GPU 操作是异步的——你调用一个 kernel,它会立刻返回,GPU 在后台执行。 所以计时的时候一定要 torch.cuda.synchronize() 等 GPU 完成,否则测出来的时间不对。

6.4.2 Profiling(性能分析)

Profiling 告诉你时间花在哪里——哪个 kernel 最慢,每个 kernel 花了多少时间。

PyTorch 自带一个 profiler:torch.profiler。

更专业的工具是 NVIDIA 的 Nsight,可以看到更详细的信息,比如: - 每个 kernel 的指令数 - 内存带宽利用率 - 占用率 - Bank 冲突

作业 2 里会用到 Nsight。

💡 优化的正确流程: 1. 先 benchmark,知道当前速度 2. 再 profile,知道哪里慢 3. 针对性优化 4. 再 benchmark,看优化效果 5. 重复

不要盲目优化——先测量,再优化。


6.5 Triton 入门:从简单例子开始

现在我们来学 Triton 编程,从最简单的例子开始。

6.5.1 例子一:GeLU(逐元素操作)

GeLU 是一个激活函数,公式是:

1
GELU(x) = 0.5 * x * (1 + tanh(sqrt(2/π) * (x + 0.044715 * x³)))

这是一个逐元素操作——每个元素独立计算,很适合 GPU 并行。

Triton 代码的结构:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
@triton.jit
def gelu_kernel(x_ptr, y_ptr, num_elements, BLOCK_SIZE: tl.constexpr):
# 1. 计算当前块处理的元素范围
pid = tl.program_id(axis=0) # 当前块的编号
start = pid * BLOCK_SIZE # 起始位置
offsets = start + tl.arange(0, BLOCK_SIZE) # 所有元素的偏移

# 2. 掩码:防止越界
mask = offsets < num_elements

# 3. 从全局内存加载数据
x = tl.load(x_ptr + offsets, mask=mask)

# 4. 计算
a = 0.79788456 * (x + 0.044715 * x * x * x)
exp = tl.exp(2 * a)
tanh = (exp - 1) / (exp + 1)
y = 0.5 * x * (1 + tanh)

# 5. 写回全局内存
tl.store(y_ptr + offsets, y, mask=mask)

解释一下每一步:

  1. 计算偏移:每个块处理 BLOCK_SIZE 个元素,块编号 pid 决定了处理哪一段
  2. 掩码(mask):如果元素总数不是 BLOCK_SIZE 的整数倍,最后一个块会越界,用 mask 来屏蔽超出的部分
  3. 加载:从全局内存读数据到寄存器
  4. 计算:在寄存器里做计算,这部分很快
  5. 存储:把结果写回全局内存

怎么调用这个 kernel?

1
2
3
4
5
6
7
def triton_gelu(x):
y = torch.empty_like(x)
num_elements = x.numel()
BLOCK_SIZE = 1024
num_blocks = triton.cdiv(num_elements, BLOCK_SIZE) # 向上取整
gelu_kernel[(num_blocks,)](x, y, num_elements, BLOCK_SIZE=BLOCK_SIZE)
return y

[(num_blocks,)] 是网格的大小——这里是一维网格,有 num_blocks 个块。

💡 逐元素操作的特点: - 最简单,每个线程处理一个(或几个)元素 - 没有线程间通信,不需要共享内存 - 通常是内存受限的——计算量小,内存访问多 - 优化重点:内存合并、算子融合

6.5.2 例子二:Softmax(归约操作)

接下来看一个稍微复杂的例子:Softmax。

Softmax 是逐行的——对矩阵的每一行做 softmax。

朴素实现的问题:

朴素的 PyTorch 实现需要多次读写内存: 1. 求每行的最大值(读一次) 2. 减去最大值(读一次,写一次) 3. 求指数(读一次,写一次) 4. 求和(读一次) 5. 归一化(读一次,写一次)

总共 5 次读、3 次写!

Triton 融合实现:

Triton 可以把所有操作融合在一个 kernel 里,只需要 1 次读、1 次写。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
@triton.jit
def softmax_kernel(x_ptr, y_ptr, x_row_stride, y_row_stride, num_cols, BLOCK_SIZE: tl.constexpr):
# 每个块处理一行
row_idx = tl.program_id(0)
col_offsets = tl.arange(0, BLOCK_SIZE)

# 加载一整行
x_start_ptr = x_ptr + row_idx * x_row_stride
x_row = tl.load(x_start_ptr + col_offsets, mask=col_offsets < num_cols, other=float("-inf"))

# 计算 softmax(全部在寄存器/共享内存里完成)
x_row = x_row - tl.max(x_row, axis=0) # 减去最大值
numerator = tl.exp(x_row)
denominator = tl.sum(numerator, axis=0)
y_row = numerator / denominator

# 写回
y_start_ptr = y_ptr + row_idx * y_row_stride
tl.store(y_start_ptr + col_offsets, y_row, mask=col_offsets < num_cols)

关键点: - 每个块处理一行 - 一整行数据加载到寄存器后,所有计算都在寄存器里完成 - 不需要反复读写全局内存 - 速度比朴素实现快好几倍!

💡 归约操作的特点: - 需要线程间协作(求 max、sum 等) - 如果数据量不大,一整行能放进一个块,就很简单 - 如果数据量大,放不下,就需要分块 + 累积

6.5.3 例子三:行求和(分块归约)

如果一行特别长,一个块放不下怎么办?

这时候就需要分块归约——把一行分成多块,每个块算一部分,最后再合并。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
@triton.jit
def row_sum_kernel(x_ptr, out_ptr, N, BLOCK_SIZE: tl.constexpr):
row = tl.program_id(0)

# 累加器
acc = tl.zeros([BLOCK_SIZE], dtype=tl.float32)

# 循环遍历所有块
for start in range(0, N, BLOCK_SIZE):
cols = start + tl.arange(0, BLOCK_SIZE)
mask = cols < N
x = tl.load(x_ptr + row * N + cols, mask=mask, other=0.0)
acc += x

# 最终归约:从 BLOCK_SIZE 个元素归约成一个标量
result = tl.sum(acc, axis=0)

tl.store(out_ptr + row, result)

思路: 1. 每个线程维护一个累加器 2. 循环遍历所有块,每块加一次 3. 最后做一次最终归约,把所有线程的结果加起来

这就是"分块归约"的思路——大的归约拆成小的,最后合并。

💡 为什么要分块? 因为一个块的大小是有限的(受共享内存、寄存器等限制),如果数据太大,一个块装不下,就必须分多次处理。


6.6 进阶例子:分块矩阵乘法

最后来看一个最复杂也最重要的例子:分块矩阵乘法。

矩阵乘法是深度学习的核心,也是 Triton 编程的"期末考试"。

6.6.1 朴素矩阵乘法的问题

朴素的矩阵乘法: - 每个元素 C[i,j] = sum_k A[i,k] * B[k,j] - 每个元素都要读 A 的一行和 B 的一列 - 内存访问量很大,算术强度很低

如果矩阵很大,这会非常慢。

6.6.2 分块的思路

分块(Tiling)的思路我们上一讲学过: 1. 把输出矩阵 C 分成小块(tile) 2. 每个小块由一个线程块负责 3. 计算时,把 A 和 B 的对应小块加载到共享内存 4. 在共享内存里做计算,结果写回全局内存

这样 A 和 B 的数据可以重复利用,大大减少全局内存访问。

6.6.3 Triton 实现的大致结构

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
@triton.jit
def matmul_kernel(a_ptr, b_ptr, c_ptr, M, N, K, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
# 每个块计算 C 的一个 BLOCK_M × BLOCK_N 的小块
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)

# 计算当前块的偏移
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, BLOCK_K)

# 累加器
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

# 循环 K 维度
for k in range(0, K, BLOCK_K):
# 加载 A 的一小块 (BLOCK_M × BLOCK_K)
a = tl.load(a_ptr + offs_m[:, None] * K + k + offs_k[None, :],
mask=(offs_m[:, None] < M) & (k + offs_k[None, :] < K),
other=0.0)

# 加载 B 的一小块 (BLOCK_K × BLOCK_N)
b = tl.load(b_ptr + (k + offs_k[:, None]) * N + offs_n[None, :],
mask=(k + offs_k[:, None] < K) & (offs_n[None, :] < N),
other=0.0)

# 矩阵乘法(在寄存器/共享内存里完成)
acc += tl.dot(a, b)

# 写回 C
c = acc.to(tl.float16)
tl.store(c_ptr + offs_m[:, None] * N + offs_n[None, :],
c,
mask=(offs_m[:, None] < M) & (offs_n[None, :] < N))

解释一下:

  1. 二维网格:矩阵乘法的网格是二维的——pid_m 是行方向的块编号,pid_n 是列方向的块编号
  2. 累加器:acc 是 BLOCK_M × BLOCK_N 的矩阵,用来累积结果
  3. K 维度循环:沿着 K 维度循环,每次加载 A 和 B 的一小块,做矩阵乘法,累加到 acc
  4. tl.dot:Triton 提供的矩阵乘法指令,会自动用 Tensor Core 加速
  5. 写回:最后把结果写回全局内存

为什么分块更快? - A 的每个 BLOCK_M × BLOCK_K 小块被 BLOCK_N 次复用 - B 的每个 BLOCK_K × BLOCK_N 小块被 BLOCK_M 次复用 - 全局内存访问量减少了约 min(BLOCK_M, BLOCK_N, BLOCK_K) 倍 - 算术强度大大提高

💡 分块大小怎么选? 分块大小是个重要的超参数,需要权衡: - 块太大:占用率低,可能放不下 - 块太小:数据复用率低,效率不高 - 常见的大小:64×64、128×128、64×128 等 - 最佳大小取决于硬件和矩阵尺寸,通常需要调参

6.6.4 额外优化:算子融合

分块矩阵乘法还有个好处——可以和其他操作融合。

比如,你想计算 GeLU(A @ B),可以把 GeLU 融合到矩阵乘法的 kernel 里:

1
2
3
4
5
# 矩阵乘法
acc += tl.dot(a, b)

# 直接在 kernel 里做 GeLU
acc = gelu(acc) # 假设 gelu 是一个逐元素函数

这样就不需要先把矩阵乘法的结果写回内存,再读出来做 GeLU——省了一次读写。

这就是算子融合的威力!


6.7 本讲小结

这一讲我们学了 Triton 编程和性能调优:

GPU 硬件细节 - Warp 和控制分支 - Warp 占用率(寄存器、共享内存的限制) - Bank 冲突(共享内存的坑) - 内存合并(全局内存的访问模式) - 块占用率(波量化问题)

性能测量 - Benchmarking:测端到端时间 - Profiling:看时间花在哪里 - 正确流程:先测量,再优化

Triton 编程 - 逐元素操作(GeLU):最简单,每个线程处理一个元素 - 归约操作(Softmax):线程间协作,一行一个块 - 分块归约(Row Sum):大归约拆成小的 - 分块矩阵乘法:最复杂也最重要,用分块减少内存访问

核心思想 - 减少全局内存访问,增加数据复用 - 算子融合,减少中间结果的读写 - 充分利用共享内存和寄存器


6.8 小白常见问题 Q&A

Q1:Triton 和 CUDA,我应该学哪个?

A:如果你是初学者,建议先学 Triton,因为: 1. 更简单,更容易上手 2. 足够强大,能写大部分高性能 kernel 3. 越来越流行,PyTorch 2.0 的编译器也用 Triton

CUDA 更底层,控制力更强,但学习曲线更陡。如果你需要做极致的优化,或者研究新的硬件特性,可以再学 CUDA。

Q2:写 Triton kernel 难吗?

A:简单的 kernel(比如逐元素操作)很容易,几行代码就搞定。

复杂的 kernel(比如分块矩阵乘法、FlashAttention)比较难,需要理解硬件特性、内存模式、分块策略等。

但好消息是——大部分时候你不需要自己写 kernel,PyTorch 自带的已经够快了。只有在特殊需求下才需要自己写。

Q3:怎么知道我的 kernel 够不够快?

A:几个方法: 1. 和 PyTorch 自带的实现比——如果比 PyTorch 慢,那肯定还有优化空间 2. 算 Roofline——看理论峰值,你的实现达到了多少 3. 用 profiler 看——内存带宽利用率、计算利用率怎么样 4. 调参试试——改改块大小,看速度怎么变

Q4:Triton 能在 AMD GPU 上跑吗?

A:目前 Triton 主要支持 NVIDIA GPU。AMD GPU 有自己的 ROCm 平台,也有类似的工具(比如 Triton 的 AMD 版本),但生态不如 NVIDIA 成熟。

这门课用的是 NVIDIA GPU,所以我们只讲 NVIDIA 的情况。


6.9 学习路线图

必须掌握: - Triton 的基本编程模型(Block、Grid) - 逐元素操作的写法 - 归约操作的基本思路 - 分块矩阵乘法的核心思想 - Benchmarking 和 Profiling 的作用

了解即可: - 具体的 Triton API 细节(用的时候查文档) - Bank 冲突的详细优化方法 - 各种高级优化技巧

配合作业: - 作业 2(Systems)里会让你写 Triton kernel - RMSNorm(逐元素 + 归约) - Flash Attention(分块 + 融合) - 可以结合这一讲的内容来做


第七讲:多 GPU 并行训练基础

7.0 前置知识:为什么需要多张 GPU?

前面几讲我们学了怎么让一张 GPU 跑得更快。但训练大模型的时候,一张 GPU 往往不够用。

为什么?两个原因:

原因一:模型太大,一张 GPU 装不下

现在的大模型动辄几十亿、几百亿参数。比如: - 7B 模型:约 14 GB(BF16) - 70B 模型:约 140 GB(BF16) - 175B 模型:约 350 GB(BF16)

而一张 A100 只有 80 GB 显存,连 70B 模型的参数都装不下,更别说还要存梯度、优化器状态、激活值了。

原因二:训练太慢,等不起

就算模型能装下,训练速度也可能太慢。比如用一张 GPU 训练一个大模型可能需要几个月甚至几年——没人等得起。

所以,我们需要多张 GPU 一起工作,把内存和计算压力分散开。

💡 生活类比: 想象你要建一座大楼: - 一张 GPU = 一个工人,既慢又搬不动太重的材料 - 多张 GPU = 一个施工队,人多力量大,分工合作

但人多了也有问题——怎么分工?怎么协调?这就是并行训练要解决的问题。


7.1 内存层次的扩展:从单卡到多卡

第二讲学过单 GPU 的内存层次,现在我们把它扩展到多 GPU、多机器的场景:

1
2
3
4
5
6
7
8
9
10
速度最快 ↑
寄存器 ← 每个线程自己的
共享内存/L1 ← 每个 SM 自己的
L2 缓存 ← 每个 GPU 自己的
HBM(显存) ← 每个 GPU 自己的
─────────────────
NVLink ← 同一台机器的 GPU 之间
─────────────────
Infiniband ← 不同机器之间
速度最慢 ↓

越往上越快,但容量越小;越往下越慢,但容量越大。

整个系统的设计原则和单 GPU 是一样的:尽量用快速的内存,减少慢速内存的访问。

只不过现在多了几层——GPU 之间的通信(NVLink)和机器之间的通信(Infiniband)。

💡 统一的思想: 从单 GPU 到多 GPU 到多机器,核心思想是一样的: - 数据局部性:尽量用离得近的、快的内存 - 减少通信:尽量少在慢的层之间搬数据 - 并行计算:尽量让多个计算单元同时工作

只是规模变大了,层次变多了。


7.2 集体操作:分布式编程的基本积木

多 GPU 编程的核心是通信——GPU 之间需要交换数据。

最基本的通信模式叫集体操作(Collective Operations)。这些是分布式编程的"积木",更复杂的并行策略都是用这些积木搭出来的。

我们一个个来看。

7.2.1 基础操作:Broadcast、Scatter、Gather、Reduce

先看四个最基础的操作。

假设我们有 4 个 GPU(rank 0 到 rank 3)。

1. Broadcast(广播)

从一个 GPU(通常是 rank 0)把数据复制到所有 GPU。

1
2
3
4
5
6
7
8
9
10
11
输入:
rank 0: [a, b, c, d]
rank 1: []
rank 2: []
rank 3: []

输出:
rank 0: [a, b, c, d]
rank 1: [a, b, c, d]
rank 2: [a, b, c, d]
rank 3: [a, b, c, d]

用途:比如初始化时,rank 0 加载模型参数,然后广播给所有 GPU。

💡 生活类比: 老师(rank 0)把讲义发给全班同学(所有 rank),每个人都拿到一份完整的讲义。

2. Scatter(分发)

把 rank 0 的数据分成 N 份,每个 GPU 拿一份。

1
2
3
4
5
6
7
8
9
10
11
输入:
rank 0: [a, b, c, d]
rank 1: []
rank 2: []
rank 3: []

输出:
rank 0: [a]
rank 1: [b]
rank 2: [c]
rank 3: [d]

用途:把数据分片,每个 GPU 负责一部分。

💡 生活类比: 班长(rank 0)把一摞作业分成 4 份,每个小组(rank)拿一份去批改。

3. Gather(收集)

Scatter 的反向操作——把所有 GPU 的数据收集到 rank 0。

1
2
3
4
5
6
7
8
9
10
11
输入:
rank 0: [a]
rank 1: [b]
rank 2: [c]
rank 3: [d]

输出:
rank 0: [a, b, c, d]
rank 1: [b]
rank 2: [c]
rank 3: [d]

用途:把各个 GPU 的计算结果汇总到一个地方。

💡 生活类比: 各小组批改完作业,把作业交回给班长(rank 0),班长把它们合在一起。

4. Reduce(归约)

所有 GPU 的数据做一个运算(比如求和、求最大值),结果放到 rank 0。

1
2
3
4
5
6
7
8
9
10
11
输入:
rank 0: [0]
rank 1: [1]
rank 2: [2]
rank 3: [3]

输出(sum):
rank 0: [6] (0+1+2+3)
rank 1: [1]
rank 2: [2]
rank 3: [3]

运算可以是 sum、min、max 等,只要满足结合律和交换律就行。

用途:比如汇总各个 GPU 的梯度。

💡 生活类比: 每个小组统计自己的捐款金额,班长把所有小组的金额加起来,得到总捐款数。


7.2.2 进阶操作:All-Gather、Reduce-Scatter、All-Reduce

上面四个操作的结果都只在 rank 0,其他 GPU 没有。

但有时候我们希望所有 GPU 都有结果——这就是 "All-" 系列操作。

1. All-Gather(全收集)

Gather 的进阶版——所有 GPU 都收集到完整的数据。

1
2
3
4
5
6
7
8
9
10
11
输入:
rank 0: [a]
rank 1: [b]
rank 2: [c]
rank 3: [d]

输出:
rank 0: [a, b, c, d]
rank 1: [a, b, c, d]
rank 2: [a, b, c, d]
rank 3: [a, b, c, d]

相当于每个 GPU 都做了一次 gather。

用途:比如每个 GPU 有一部分参数,需要把所有参数合起来才能做前向传播。

💡 生活类比: 每个小组有一部分拼图,大家把自己的部分传出去,最后每个人都有完整的拼图。

2. Reduce-Scatter(归约-分发)

先 reduce,再 scatter。

每个维度先求和,然后把结果分发到各个 GPU。

1
2
3
4
5
6
7
8
9
10
11
输入:
rank 0: [0, 1, 2, 3]
rank 1: [1, 2, 3, 4]
rank 2: [2, 3, 4, 5]
rank 3: [3, 4, 5, 6]

输出(sum):
rank 0: [6] (0+1+2+3) 第 0 维的和
rank 1: [10] (1+2+3+4) 第 1 维的和
rank 2: [14] (2+3+4+5) 第 2 维的和
rank 3: [18] (3+4+5+6) 第 3 维的和

用途:比如梯度求和后,每个 GPU 只负责更新自己那部分参数。

💡 生活类比: 每个小组有一张成绩表(各科分数),大家先把各科的总分算出来(reduce),然后每个小组负责保管一科的成绩(scatter)。

3. All-Reduce(全归约)

最重要的操作!几乎所有分布式训练都会用到。

All-Reduce = Reduce-Scatter + All-Gather

先 reduce-scatter,再 all-gather,最后每个 GPU 都有完整的归约结果。

1
2
3
4
5
6
7
8
9
10
11
输入:
rank 0: [0, 1, 2, 3]
rank 1: [1, 2, 3, 4]
rank 2: [2, 3, 4, 5]
rank 3: [3, 4, 5, 6]

输出(sum):
rank 0: [6, 10, 14, 18]
rank 1: [6, 10, 14, 18]
rank 2: [6, 10, 14, 18]
rank 3: [6, 10, 14, 18]

所有 GPU 的数据都被求和了,而且每个 GPU 都有完整的结果。

用途:数据并行训练中,各个 GPU 算出自己的梯度后,用 all-reduce 把梯度汇总,这样每个 GPU 都有完整的梯度,可以同步更新参数。

💡 为什么要拆成两步? 你可能会问:既然 all-reduce = reduce-scatter + all-gather,为什么不直接做 all-reduce?

因为拆成两步更灵活!比如 ZeRO(后面会讲)就利用了这一点——只做 reduce-scatter,不做 all-gather,这样可以节省内存。

在带宽受限的情况下,all-reduce 的通信量是 2×数据量(reduce-scatter 1×,all-gather 1×),这是理论最优的。


7.2.3 最灵活的操作:All-to-All

最后一个,也是最复杂的:All-to-All。

每个 GPU 给其他每个 GPU 发一些数据。

1
2
3
4
5
6
7
8
9
10
11
输入:
rank 0: [a0, a1, a2, a3] → a0 给自己, a1 给 rank1, a2 给 rank2, a3 给 rank3
rank 1: [b0, b1, b2, b3] → b0 给 rank0, b1 给自己, b2 给 rank2, b3 给 rank3
rank 2: [c0, c1, c2, c3] → c0 给 rank0, c1 给 rank1, c2 给自己, c3 给 rank3
rank 3: [d0, d1, d2, d3] → d0 给 rank0, d1 给 rank1, d2 给 rank2, d3 给自己

输出:
rank 0: [a0, b0, c0, d0] ← 从所有 rank 收到的第 0 份
rank 1: [a1, b1, c1, d1] ← 从所有 rank 收到的第 1 份
rank 2: [a2, b2, c2, d2] ← 从所有 rank 收到的第 2 份
rank 3: [a3, b3, c3, d3] ← 从所有 rank 收到的第 3 份

看起来像是矩阵的转置——输入是按行存的,输出是按列存的。

用途:MoE 模型的路由——每个 GPU 有一部分数据和一部分专家,需要把数据送到对应专家的 GPU 上。

💡 生活类比: 想象一个邮局系统: - 每个城市(rank)都有一堆信件,分别寄往不同的城市 - All-to-all 就是把所有信件按目的地分类,每个城市收到寄给自己的所有信件

这是最复杂的通信模式,但也是最灵活的。


7.3 硬件:GPU 之间怎么连?

知道了通信操作,我们来看看硬件层面 GPU 之间是怎么连接的。

同一台机器里的 GPU 怎么通信?

PCIe:最传统的方式。GPU 插在 PCIe 插槽上,通过 PCIe 总线通信。 - PCIe 4.0 x16:约 32 GB/s - 不算特别快,但够用

NVLink / NVSwitch:NVIDIA 的高速互联技术。 - NVLink:GPU 之间直接连,比 PCIe 快很多 - NVSwitch:像一个交换机,所有 GPU 都连上去,全互联 - H100 的 NVLink 4.0:约 900 GB/s - B200 的 NVLink 5.0:约 1.8 TB/s

NVLink 比 PCIe 快几十倍!所以同一台机器里的 GPU 通信很快。

💡 生活类比: - PCIe = 普通公路,车多了就堵 - NVLink = 高速公路,又宽又快 - NVSwitch = 立交桥,各个方向都能快速到达

7.3.2 节点间:Ethernet → Infiniband

不同机器之间怎么通信?

Ethernet(以太网):最常见的网络。 - 普通的 10G/25G 以太网 - 便宜,但慢,延迟高 - 需要经过 CPU(TCP/IP 协议栈)

Infiniband(无限带宽):高性能计算常用的网络。 - 速度快(200G/400G) - 延迟低 - 支持 RDMA(Remote Direct Memory Access):一个 GPU 可以直接读写另一个机器的 GPU 内存,不需要经过 CPU

RDMA 很重要——如果没有 RDMA,数据要先从 GPU 搬到 CPU,再通过网络发出去,到了另一边再从 CPU 搬到 GPU,既慢又占 CPU。

有了 RDMA,GPU 之间可以直接通信,绕过 CPU,快很多。

💡 生活类比: - 普通以太网 = 普通快递,要经过很多中转站,慢 - Infiniband + RDMA = 直达快递,从发货地直接送到收货地,中间不停,快

RDMA 就是"直达"的意思。

7.3.3 NCCL:通信的"翻译官"

这么多硬件、这么多通信模式,程序员怎么用?

不用担心,NVIDIA 提供了一个库叫 NCCL(NVIDIA Collective Communications Library)。

NCCL 做什么? - 自动检测硬件拓扑(有多少 GPU、怎么连的、网络是什么) - 自动选择最优的通信路径 - 提供统一的接口(all-reduce、all-gather 等)

你只需要调用 NCCL 的 all_reduce 函数,它会自动用最快的方式完成通信——不管是 NVLink 还是 Infiniband。

PyTorch 的分布式训练底层就是用 NCCL 的。

💡 生活类比: NCCL 就像一个物流系统——你只需要告诉它"把这个包裹送到那个地方",它自动选择最快的路线(公路、铁路、飞机),你不用管中间怎么运。


7.4 分布式训练策略:三种基本方式

有了通信的积木,我们来看怎么用它们来实现分布式训练。

主要有三种基本的并行策略:

  1. 数据并行(Data Parallelism):切 batch
  2. 张量并行(Tensor Parallelism):切矩阵的维度
  3. 流水线并行(Pipeline Parallelism):切层

我们一个个来看。


7.4.1 数据并行(Data Parallelism / DDP)

最简单、最常用的并行方式。

核心思想:每个 GPU 都有完整的模型,但是输入数据被分成 N 份,每个 GPU 处理 1/N 的数据。

步骤: 1. 把 batch 分成 N 份,每个 GPU 拿一份 2. 每个 GPU 用自己的那份数据做前向传播和反向传播,算出梯度 3. 用 all-reduce 把所有 GPU 的梯度汇总(求和) 4. 每个 GPU 用汇总后的梯度更新自己的参数

因为所有 GPU 的初始参数一样,梯度也一样,所以更新后的参数也一样——所有 GPU 的模型始终保持同步。

通信量:每个 step 通信 2×参数量(all-reduce 的通信量)

内存:每个 GPU 都要有完整的模型 → 没有内存扩展

计算:线性扩展,N 个 GPU 就是 N 倍的计算速度(理想情况下)

💡 生活类比: 想象一个工厂生产产品: - 每个车间(GPU)都有完整的生产线(完整模型) - 原材料(batch)被分成几份,每个车间加工一份 - 加工完后,大家把质检报告(梯度)汇总一下(all-reduce) - 每个车间根据汇总的报告调整生产线(参数更新)

优点是简单,缺点是每个车间都要有一整条生产线(内存没省)。

适用场景: - 模型不大,单卡能装下 - 想加快训练速度 - 最常用、最稳定的并行方式

PyTorch 里的 DistributedDataParallel(DDP) 就是数据并行。


7.4.2 张量并行(Tensor Parallelism)

第二种方式:张量并行,也叫模型并行的一种。

核心思想:把矩阵乘法的维度切开,每个 GPU 只算一部分。

比如矩阵乘法 Y = XA: - 把 A 按列切成 N 份:[A1, A2, ..., AN] - 每个 GPU 算 X × Ai,得到 Yi - 最后把所有 Yi 拼起来,得到 Y = [Y1, Y2, ..., YN]

这样每个 GPU 只存 1/N 的参数,计算量也是 1/N。

但是——每次矩阵乘法都需要通信!因为: - 前向传播:算完后需要 all-gather 把结果拼起来 - 反向传播:梯度需要 all-reduce

通信量:每层都要通信 → 通信量很大

内存:参数分片,每个 GPU 只存 1/N → 内存扩展 N 倍

计算:线性扩展

💡 生活类比: 想象一个大的计算任务: - 把一个大任务拆成几个小任务,每个人(GPU)做一部分 - 做完后大家把结果拼起来 - 但每一步都要同步,因为下一步依赖上一步的结果

优点是省内存,缺点是通信多,对网络速度要求高。

适用场景: - 模型很大,单卡装不下 - 网络很快(比如 NVLink) - 通常和数据并行一起用

Megatron-LM 就是张量并行的代表实现。

⚠️ 注意:张量并行对网络速度要求很高,因为每层都要通信。 通常只在同一台机器内(NVLink 连接)做张量并行,跨机器的话通信太慢了。


7.4.3 流水线并行(Pipeline Parallelism)

第三种方式:流水线并行。

核心思想:把模型的层切开,不同的 GPU 跑不同的层。

比如一个 8 层的模型,4 个 GPU: - GPU 0:第 1-2 层 - GPU 1:第 3-4 层 - GPU 2:第 5-6 层 - GPU 3:第 7-8 层

数据从 GPU 0 流进去,经过 GPU 1、GPU 2,最后从 GPU 3 出来。

就像工厂的流水线一样——每个工人只做一道工序,产品从一头进去,另一头出来。

问题:流水线气泡(Pipeline Bubble)

朴素的流水线有个问题:GPU 利用率不高。

比如一个 batch 的数据: - 时间 1:GPU 0 在算,其他 GPU 闲着 - 时间 2:GPU 1 在算,其他 GPU 闲着 - 时间 3:GPU 2 在算,其他 GPU 闲着 - 时间 4:GPU 3 在算,其他 GPU 闲着

每个时刻只有一个 GPU 在工作,其他都在等——这叫流水线气泡。

怎么减少气泡?

把 batch 切成更小的 micro-batch,让流水线"填满"。

比如把 batch 切成 8 个 micro-batch: - 时间 1:GPU 0 算 micro-batch 0 - 时间 2:GPU 0 算 micro-batch 1,GPU 1 算 micro-batch 0 - 时间 3:GPU 0 算 micro-batch 2,GPU 1 算 micro-batch 1,GPU 2 算 micro-batch 0 - ... - 时间 N:所有 GPU 都在算(流水线满了)

这样 GPU 利用率就高多了。

但气泡还是存在的——流水线启动和结束的时候还是有空隙。micro-batch 越多,气泡的比例越小,但调度开销越大。

💡 生活类比: 工厂的流水线: - 如果只有一个产品,那每个时刻只有一个工人在干活,其他人等着(气泡大) - 如果有很多产品,一个接一个流过去,每个工人都不停干活(气泡小) - 但启动的时候(第一个产品还没流到后面)和结束的时候(最后一个产品流走了)还是有空隙

通信量:每个 micro-batch 在层之间传一次激活值 → 通信量中等

内存:每个 GPU 只存 1/N 的层 → 内存扩展 N 倍

计算:有气泡,效率低于线性扩展

适用场景: - 模型非常深,单卡装不下 - 可以和数据并行、张量并行组合使用 - 对网络速度的要求介于数据并行和张量并行之间


7.4.4 三种方式的对比

并行方式 切什么 通信量 内存扩展 计算扩展 网络要求
数据并行 batch 每个 step 2×参数 无 线性 低
张量并行 矩阵维度 每层都通信 N 倍 线性 很高
流水线并行 层 每个 micro-batch 传激活 N 倍 接近线性(有气泡) 中等

实际中怎么用?

大模型训练通常是三种方式组合使用:

比如 256 个 GPU 训练一个大模型: - 8 个 GPU 做张量并行(同一台机器,NVLink) - 4 个 stage 做流水线并行 - 8 个做数据并行 - 总共 8 × 4 × 8 = 256 个 GPU

这样既省内存,又有速度,还能扩展到很多 GPU。

💡 记忆口诀: - 数据并行:切数据,每个 GPU 有完整模型 - 张量并行:切矩阵,每个 GPU 有部分参数 - 流水线并行:切层,每个 GPU 有几层 - 实际中:三种一起用


7.5 本讲小结

这一讲我们学了多 GPU 并行训练的基础:

通信基础 - 集体操作:Broadcast、Scatter、Gather、Reduce - 进阶操作:All-Gather、Reduce-Scatter、All-Reduce - 最灵活:All-to-All(MoE 用)

硬件 - 节点内:PCIe → NVLink/NVSwitch(快) - 节点间:Ethernet → Infiniband + RDMA(较慢) - NCCL:自动优化通信路径

三种并行策略 - 数据并行(DDP):切 batch,简单常用,无内存扩展 - 张量并行:切矩阵维度,省内存,通信多 - 流水线并行:切层,省内存,有气泡 - 实际中组合使用


7.6 小白常见问题 Q&A

Q1:数据并行和 ZeRO/FSDP 是什么关系?

A:ZeRO/FSDP 是数据并行的进阶版本。

普通的数据并行(DDP):每个 GPU 都有完整的参数、梯度、优化器状态 → 内存浪费。

ZeRO/FSDP:把参数、梯度、优化器状态也分片,每个 GPU 只存一部分 → 节省内存。

下一讲我们会详细讲 ZeRO。

Q2:为什么张量并行通常只在 8 卡以内做?

A:因为张量并行的通信量很大,每层都要通信。

如果跨机器(用 Infiniband),通信速度比 NVLink 慢很多,张量并行的开销会非常大。

所以通常的做法是: - 同一台机器的 8 张卡(NVLink 连接)做张量并行 - 跨机器做数据并行或流水线并行

Q3:流水线并行的气泡能完全消除吗?

A:不能完全消除,但可以尽量减小。

方法: 1. 增加 micro-batch 数量 → 气泡比例减小,但调度开销增加 2. 用更聪明的调度策略(比如 GPipe、PipeDream) 3. 把气泡和计算重叠(比如反向传播的气泡和前向传播的计算重叠)

但不管怎样,启动和结束的时候总有一些空隙,气泡不可能完全为零。

Q4:训练大模型时,三种并行的比例怎么选?

A:这是个复杂的问题,取决于: - 模型大小(越大越需要模型并行) - 机器数量(越多越需要数据并行) - 网络速度(网络越快,越能做模型并行) - 内存预算

一般的经验法则: - 张量并行:2-8(同一台机器) - 流水线并行:2-16 - 数据并行:剩下的

具体需要做性能测试来找到最优配置。


7.7 学习路线图

必须掌握: - 集体操作的概念(特别是 All-Reduce) - 数据并行的工作原理 - 三种并行方式的区别和各自的优缺点 - NVLink 和 Infiniband 的区别

了解即可: - 各种集体操作的具体实现 - 流水线并行的详细调度算法 - NCCL 的内部工作原理

配合作业: - 作业 2(Systems)里有分布式训练的内容,可以结合这一讲理解 - 下一讲(L8)会讲 ZeRO 和更详细的并行策略,可以接着看


第八讲:并行训练进阶 — ZeRO 与大规模训练

8.0 前置知识:数据并行的内存问题

上一讲学了数据并行(DDP),它简单、稳定、容易用。但它有个大问题:内存没有扩展。

每个 GPU 都要存完整的模型——参数、梯度、优化器状态,全都要存一份。

我们来算笔账:训练一个模型,每个参数需要多少内存?

  • 模型参数(BF16):2 字节
  • 梯度(BF16):2 字节
  • 优化器状态(Adam):
    • FP32 主权重:4 字节
    • 一阶矩估计(m):4 字节
    • 二阶矩估计(v):4 字节

加起来:2 + 2 + 4 + 4 + 4 = 16 字节/参数!

也就是说,一个 7B 参数的模型,训练时需要约 112 GB 显存——一张 A100(80 GB)都装不下!

更别说 70B、175B 的模型了,根本装不下。

那怎么办?这就是 ZeRO 要解决的问题。

💡 生活类比: 想象一个团队做项目,每个人都要有一份完整的项目资料(参数、梯度、优化器状态)。 - 人少的时候还好,资料能放下 - 人多了,每个人都存一份,太浪费空间了

能不能大家分工,每个人只存一部分,需要的时候再互相传?这就是 ZeRO 的思路。


8.1 ZeRO:零冗余优化

ZeRO(Zero Redundancy Optimizer) 是微软在 2020 年提出的技术,核心思想很简单:

数据并行中,每个 GPU 存的东西有大量重复。我们把这些重复的去掉,每个 GPU 只存一部分。

具体来说,ZeRO 分三个阶段,分片的程度越来越深:

阶段 分片什么 每个 GPU 存多少
ZeRO Stage 1 优化器状态 优化器状态 1/N,参数和梯度完整
ZeRO Stage 2 优化器状态 + 梯度 优化器状态和梯度各 1/N,参数完整
ZeRO Stage 3 优化器状态 + 梯度 + 参数 所有东西各 1/N

N 是 GPU 数量。

我们一个个来看。


8.1.1 ZeRO Stage 1:优化器状态分片

最基础的 ZeRO:只把优化器状态分片。

内存节省: - 优化器状态:原来 12 字节/参数(FP32 主权重 + m + v),现在 12/N 字节/参数 - 参数和梯度:还是完整的,各 2 字节/参数

总内存:4 + 12/N 字节/参数

比如 N=8: - 原来:16 字节/参数 - 现在:4 + 12/8 = 5.5 字节/参数 - 节省了约 65% 的内存!

怎么工作的?

步骤: 1. 前向传播:每个 GPU 有完整的参数,正常计算(和 DDP 一样) 2. 反向传播:每个 GPU 算出完整的梯度(和 DDP 一样) 3. Reduce-Scatter 梯度:把梯度分片,每个 GPU 只留自己那部分 4. 优化器更新:每个 GPU 用自己那部分梯度和优化器状态,更新自己那部分参数 5. All-Gather 参数:把更新后的参数汇总,每个 GPU 都有完整的参数

通信量: - Reduce-Scatter:1×参数量 - All-Gather:1×参数量 - 总共:2×参数量 → 和 DDP 的 all-reduce 一样!

💡 神奇的地方: ZeRO Stage 1 的通信量和 DDP 完全一样,但内存省了很多! 为什么?因为 all-reduce 本来就可以拆成 reduce-scatter + all-gather,通信量一样。 ZeRO 只是利用了这个拆分,把中间状态分片存储了。

这就是为什么说 ZeRO Stage 1 是"免费的内存收益"——通信没变,内存却省了。

💡 生活类比: 团队做项目,每个人都有完整的文档(参数),但计算结果(梯度)汇总后,每个人只负责更新自己那部分内容(优化器状态)。 更新完了,再把所有部分拼起来,每个人都拿到完整的新版本。

这样每个人不需要存所有的修改记录(优化器状态),只存自己负责的那部分就行,省空间。


8.1.2 ZeRO Stage 2:梯度也分片

Stage 1 只分片了优化器状态,Stage 2 更进一步:梯度也分片。

内存节省: - 优化器状态:12/N 字节/参数 - 梯度:2/N 字节/参数(原来 2 字节) - 参数:还是完整的,2 字节/参数

总内存:2 + 14/N 字节/参数

比如 N=8: - 原来:16 字节/参数 - 现在:2 + 14/8 = 3.75 字节/参数 - 节省了约 77%!

怎么工作的?

Stage 2 和 Stage 1 类似,但梯度不需要完整保存——反向传播时,算出一层的梯度就立刻 reduce-scatter 出去,不用存完整的梯度。

这样就不需要在内存里保留完整的梯度了,进一步省内存。

通信量:还是 2×参数量,和 Stage 1、DDP 一样。

⚠️ 注意:Stage 2 有个限制——你不能在反向传播完成后再用完整的梯度做什么事情(比如梯度裁剪),因为梯度已经被分片了。 当然,梯度裁剪可以在分片前做,或者用其他方式处理。


8.1.3 ZeRO Stage 3:参数也分片(FSDP)

Stage 3 是最激进的:连参数都分片,每个 GPU 只存 1/N 的参数。

这就是 PyTorch 里的 FSDP(Fully Sharded Data Parallel)。

内存节省: - 优化器状态:12/N 字节/参数 - 梯度:2/N 字节/参数 - 参数:2/N 字节/参数

总内存:16/N 字节/参数

比如 N=8: - 原来:16 字节/参数 - 现在:16/8 = 2 字节/参数 - 节省了 87.5%!

内存几乎是线性扩展的——N 个 GPU,每个 GPU 只需要 1/N 的内存。

怎么工作的?

Stage 3 更复杂,因为参数也是分片的,前向传播时需要先把参数合起来:

  1. 前向传播:
    • 每一层计算前,先 all-gather 这一层的参数(拿到完整参数)
    • 计算
    • 计算完后,扔掉完整参数(只保留自己的分片)
  2. 反向传播:
    • 每一层计算前,先 all-gather 这一层的参数
    • 计算梯度
    • 计算完后,reduce-scatter 梯度

通信量:比 DDP 多,因为每层都要 all-gather 参数。

具体多多少?取决于层数。如果有 L 层,那通信量大约是 2×L×参数量?不对,等一下——

实际上,Stage 3 的通信量和 DDP 一样,都是 2×参数量。为什么?

因为: - 前向传播:每层 all-gather 参数 → 总共 L 次 all-gather,每次 1/N×参数量 → 总共 1×参数量 - 反向传播:每层 reduce-scatter 梯度 → 总共 L 次 reduce-scatter,每次 1/N×参数量 → 总共 1×参数量 - 合计:2×参数量

哦,对!因为每层只传 1/N 的参数,L 层加起来正好是 1×参数量。

所以 Stage 3 的通信量也是 2×参数量,和 DDP 一样!

⚠️ 等等,这不对吧? 你可能会想:每层都要通信,那通信次数不是多了很多吗?

是的,通信次数多了,但总通信量(字节数)是一样的。

但通信次数多意味着延迟更高——因为每次通信都有开销。 所以 Stage 3 虽然总通信量没变,但延迟可能更高,速度可能更慢。

💡 生活类比: 想象你要搬一堆书到另一个房间: - DDP:一次把所有书搬过去(一次通信,量大) - ZeRO Stage 3:一本一本地搬(多次通信,每次量小)

总重量(总字节数)一样,但一本一本地搬要跑很多趟,花的时间更多(延迟更高)。

怎么减少通信次数?

可以用参数分组(parameter grouping)——把几层的参数合在一起,一次 all-gather 一组,减少通信次数。

但这样内存占用会增加,因为要同时存一组的完整参数。

这是内存和通信的权衡: - 分组大 → 通信次数少 → 速度快 → 内存占用多 - 分组小 → 通信次数多 → 速度慢 → 内存占用少


8.1.4 ZeRO 三个阶段的对比

Stage 1 Stage 2 Stage 3 (FSDP)
分片内容 优化器状态 优化器 + 梯度 优化器 + 梯度 + 参数
内存(N=8) ~5.5 字节/参数 ~3.75 字节/参数 ~2 字节/参数
通信量 2×参数 2×参数 2×参数
通信次数 1 次/step 1 次/step 多次(每层一次)
实现复杂度 简单 中等 复杂
速度 和 DDP 差不多 和 DDP 差不多 可能稍慢

怎么选? - 内存够的话,用 DDP 最简单 - 内存差一点,用 Stage 1 或 Stage 2 - 内存差很多,用 Stage 3 (FSDP) - 极致省内存,用 Stage 3 + 激活检查点 + 卸载到 CPU

💡 ZeRO-Offload 和 ZeRO-Infinity: 微软还提出了更激进的版本: - ZeRO-Offload:把优化器状态和梯度卸载到 CPU 内存 - ZeRO-Infinity:还可以卸载到磁盘(NVMe)

这样可以用更少的 GPU 显存训练更大的模型,但速度会更慢,因为要和 CPU/磁盘交换数据。


8.2 其他并行方式

除了 ZeRO,还有一些其他的并行方式。

8.2.1 序列并行(Sequence Parallelism)

序列并行是和张量并行配合使用的一种技术。

张量并行把矩阵的维度切开,但有些操作(比如 LayerNorm、Dropout)是按序列维度做的,不好切。

序列并行的思路是:沿着序列维度分片,把这些操作也分布到多个 GPU 上。

这样可以进一步减少每个 GPU 的内存占用,也能提高计算效率。

序列并行通常和张量并行一起用,比如 Megatron-LM 里就有。

💡 生活类比: 张量并行是"竖着切"(切特征维度),序列并行是"横着切"(切序列维度)。 横竖都切,每个 GPU 的工作量就更小了。


8.2.2 专家并行(Expert Parallelism)

第四讲学过 MoE(混合专家模型)。MoE 天然适合并行——每个 GPU 放几个专家,token 在 GPU 之间路由。

这就是专家并行。

专家并行的通信模式是 all-to-all——每个 GPU 把要去其他 GPU 的 token 发过去,同时接收发给自己的 token。

专家并行的好处: - 内存扩展好:每个 GPU 只存自己的专家 - 计算扩展好:每个 GPU 只算自己的专家

但也有挑战: - 负载均衡:有的 GPU 忙,有的闲 - 通信复杂:all-to-all 通信模式复杂 - 需要 MoE 模型

DeepSeek、Mixtral 等 MoE 模型训练时,都会用到专家并行。


8.2.3 数据并行 + 模型并行的组合

实际训练大模型时,通常是多种并行方式组合使用。

比如训练一个 175B 参数的模型,用 512 张 GPU:

1
512 GPU = 8 (张量并行) × 4 (流水线并行) × 16 (数据并行/FSDP)
  • 张量并行:同一台机器的 8 张卡,NVLink 高速互联
  • 流水线并行:4 台机器,每台负责几层
  • 数据并行/FSDP:16 组,每组有完整的模型(分片存储)

这样组合起来,既省内存,又有速度,还能扩展到很多 GPU。

💡 为什么要这么复杂? 因为每种并行方式都有自己的优缺点: - 数据并行:简单,但不省内存 - 张量并行:省内存,但通信量大,需要高速网络 - 流水线并行:省内存,但有气泡 - FSDP:省内存,但通信次数多

组合使用可以取长补短,达到最优的效果。


8.3 大规模训练的挑战

训练超大规模的模型(比如千亿、万亿参数),除了并行策略,还有很多其他挑战。

8.3.1 稳定性

模型越大,训练越容易不稳定。

可能出现的问题: - 梯度爆炸/消失 - 损失突然飙升(loss spike) - 某些层的数值异常

怎么应对? - 梯度裁剪(Gradient Clipping) - 学习率 warmup - 各种归一化(RMSNorm、QK-Norm、Z-Loss 等) - 混合精度训练时,某些关键操作用 FP32 - 遇到 loss spike 时,回退到之前的 checkpoint 重训

💡 生活类比: 小团队(小模型)好管理,大团队(大模型)容易出问题——一个人出错可能影响整个团队。 所以大团队需要更完善的制度(各种稳定化技巧)来保证正常运转。

8.3.2 检查点(Checkpointing)

训练大模型可能需要几周甚至几个月,中间随时可能出问题(硬件故障、网络中断、软件 bug 等)。

所以需要定期保存检查点(Checkpoint)——把模型参数、优化器状态、训练进度等都存到磁盘上。

出问题了,可以从最近的检查点恢复,不用从头开始。

但大模型的检查点也很大——一个 175B 模型的检查点可能有几百 GB 甚至几 TB。保存和加载都需要时间。

怎么优化? - 异步保存:保存检查点时不阻塞训练 - 增量保存:只保存变化的部分 - 压缩:用压缩算法减少磁盘占用


8.3.3 故障恢复

大规模训练用的 GPU 很多(几千张),硬件故障的概率就高了。

比如 1000 张 GPU,假设每张 GPU 的平均无故障时间是 1 年,那平均每天就有约 3 张 GPU 出故障。

所以训练系统必须能容错——某个 GPU 挂了,训练能继续,或者快速恢复。

怎么做? - 定期保存检查点,挂了就从检查点恢复 - 弹性训练:GPU 数量变化时能自动调整 - 预测性维护:提前发现可能出问题的 GPU


8.3.4 成本

训练大模型非常非常贵。

粗略估算: - 1000 张 H100,训练 1 个月 - 每张 H100 约 $30,000(但通常是租用,约 $2-4/小时) - 1000 张 × $3/小时 × 24 小时 × 30 天 = $2,160,000 - 也就是约 1500 万人民币,一个月!

这还只是 GPU 费用,还有电力、制冷、网络、人力等等。

所以训练大模型是个非常昂贵的事情,只有大公司才玩得起。

💡 为什么大家都在做推理优化? 因为训练虽然贵,但只是一次性的。推理是持续的——用户每次提问都要推理,用量大得多。 所以推理成本可能比训练成本高得多,优化推理更有商业价值。


8.4 TPU vs GPU:网络架构的差异

最后我们来看看 TPU 和 GPU 在网络架构上的区别,这也是两种硬件设计思路不同的体现。

8.4.1 TPU 的环形网格(Toroidal Mesh)

TPU 的网络是环形网格(toroidal mesh)结构——TPU 排成一个二维网格,每个 TPU 只和上下左右四个邻居相连。

1
2
3
4
5
6
7
┌───┬───┬───┬───┐
│ 0 │ 1 │ 2 │ 3 │
├───┼───┼───┼───┤
│ 4 │ 5 │ 6 │ 7 │
├───┼───┼───┼───┤
│ 8 │ 9 │10 │11 │
└───┴───┴───┴───┘

每个 TPU 只有 4 个连接,不是全互联。

这样设计的好处: - 布线简单,容易扩展 - 每个连接的带宽可以做得很高

坏处: - 远距离通信要经过很多跳,延迟高 - all-to-all 之类的操作效率低

适合什么?张量并行——因为张量并行的通信模式是规整的,正好匹配环形网格。

8.4.2 GPU 的全互联(All-to-All)

GPU 的网络(通过 NVSwitch 和 Infiniband)更偏向全互联——每个 GPU 都能直接和其他所有 GPU 通信。

好处: - 通信模式灵活 - all-to-all、all-reduce 之类的操作效率高 - 适合 MoE、数据并行等

坏处: - 布线复杂,扩展到大规模时成本高 - 每个连接的带宽相对较低

💡 设计哲学的差异: - TPU:为特定的计算模式(矩阵乘法、张量并行)优化,效率高但不够灵活 - GPU:更通用,什么都能做,但特定场景下效率可能不如 TPU

这就像专用芯片和通用芯片的区别——专用的更快更省电,但只能做特定的事;通用的什么都能做,但效率稍低。

8.4.3 Mesh vs Tree 拓扑

还有两种常见的网络拓扑:

Mesh(网状):每个节点和邻居相连,像一张网。 - 优点:容错性好,一条线断了还有其他路径 - 缺点:远距离通信跳数多

Tree(树状):节点组成树的结构,从上到下分层。 - 优点:适合广播、汇聚之类的操作 - 缺点:根节点是瓶颈,容错性差

实际的网络通常是混合的——比如 GPU 集群里,机柜内是全互联(通过交换机),机柜之间是树状(fat-tree)。


8.5 本讲小结

这一讲我们学了并行训练的进阶知识:

ZeRO 三阶段 - Stage 1:优化器状态分片,免费的内存收益 - Stage 2:梯度也分片,更省内存 - Stage 3 (FSDP):参数也分片,极致省内存 - 通信量都是 2×参数量,但 Stage 3 通信次数多

其他并行方式 - 序列并行:沿序列维度切,配合张量并行 - 专家并行:MoE 天然适合,all-to-all 通信 - 组合使用:实际中多种并行方式混合

大规模训练的挑战 - 训练稳定性 - 检查点与故障恢复 - 成本高昂

网络架构 - TPU:环形网格,适合张量并行 - GPU:全互联,更灵活 - Mesh vs Tree 拓扑


8.6 小白常见问题 Q&A

Q1:ZeRO 和张量并行,哪个更好?

A:它们解决的问题不同,不是替代关系,通常一起用。

  • ZeRO 是数据并行的进阶,主要解决"内存不够"的问题,实现简单,通信量不大
  • 张量并行是模型并行的一种,也解决"内存不够"的问题,但通信量大,需要高速网络

实际中,通常是: - 同一台机器内(NVLink):张量并行 - 跨机器:数据并行 + ZeRO/FSDP

Q2:FSDP 和 DeepSpeed ZeRO 是什么关系?

A:DeepSpeed 是微软的深度学习优化库,ZeRO 是 DeepSpeed 里的一个功能。

FSDP(Fully Sharded Data Parallel)是 PyTorch 官方实现的类似 ZeRO Stage 3 的功能。

它们的核心思想是一样的——把参数、梯度、优化器状态都分片。只是实现不同,API 不同。

现在 PyTorch 官方推荐用 FSDP,因为它集成在 PyTorch 里,更方便。

Q3:训练大模型,用多少 GPU 最合适?

A:这取决于很多因素: - 模型大小:越大需要的 GPU 越多 - 时间预算:越急需要的 GPU 越多 - 预算:GPU 很贵,要考虑成本 - 效率:GPU 太多的话,通信开销会增大,效率下降

一般来说,有个"缩放定律"——GPU 数量增加,训练速度不是线性增加的,因为通信开销会越来越大。

所以不是 GPU 越多越好,要找到性价比最高的点。

Q4:为什么大模型训练要用流水线并行?直接用 FSDP 不行吗?

A:FSDP 确实能省内存,但有个问题:通信次数多,延迟高。

如果模型特别大,层数特别多,FSDP 的通信开销会很大,速度会很慢。

流水线并行的好处是: - 通信量中等(只在层之间传激活值) - 可以和计算重叠(一边算一边通信) - 适合超大规模模型

所以超大模型通常会用流水线并行 + 张量并行 + 数据并行/FSDP 的组合。


8.7 学习路线图

必须掌握: - ZeRO 的核心思想(去掉冗余,分片存储) - ZeRO 三个阶段的区别 - 为什么 ZeRO 的通信量和 DDP 一样 - FSDP 是什么

了解即可: - 序列并行、专家并行的概念 - 大规模训练的各种挑战 - TPU 和 GPU 网络架构的区别 - Mesh vs Tree 拓扑

配合作业: - 作业 2(Systems)里有分布式训练的内容,可以结合理解 - 整个系统与并行模块(L4-L8)到这里就结束了,可以回头复习一下


模块小结:系统与并行(L4-L8)

这五讲我们从"一张 GPU 怎么跑得快"讲到"很多 GPU 怎么一起干活",核心思想其实是一致的:

1. 内存层次结构 - 单 GPU:寄存器 → 共享内存 → L2 → HBM - 多 GPU:HBM → NVLink → Infiniband - 原则:尽量用快的内存,减少慢的内存的访问

2. 减少通信/内存访问 - 算子融合:减少中间结果的读写 - 分块(Tiling):增加数据复用 - 各种并行策略:减少需要传输的数据量

3. 并行计算 - 单 GPU:几千个线程同时跑 - 多 GPU:很多 GPU 同时跑 - 关键:怎么分工、怎么协调

4. 权衡(Trade-off) - 内存 vs 计算:重计算用计算换内存 - 内存 vs 通信:FSDP 用通信换内存 - 速度 vs 通用性:专用硬件快但不够灵活

整个系统优化的本质,就是在各种约束条件下找到最优的平衡点。



CS336 学习笔记 02:注意力替代、MoE 与 GPU 系统
https://cdro.tech/notes/CS/cs336-02-l4-l8-attention-alternatives-moe-gpu-systems-triton/
作者
k9Q6CK42
发布于
2026年9月28日
更新于
2026年9月28日
许可协议