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 | |
先算 Q 和 K 的内积(得到一个 n×n 的大矩阵),再乘 V。这个 n×n 的大矩阵就是 O(n²) 的来源。
注意公式里要除以 √d_k(每个头的维度)做缩放:维度一大,点积的数值会随 √d_k 增大,softmax 容易饱和成“一家独大”、梯度几乎消失;除以 √d_k 把点积方差拉回 1,训练才稳定。
那能不能换个顺序,先算 K 和 V,再乘 Q 呢?
1 | |
这样就变成了: - 先算 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 | |
什么意思呢?就是你可以一个 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 | |
其中 γ_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 | |
这里多了个 β_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 路由:
- 用一个简单的线性层(叫"门控"或"路由器")给每个专家打分
- 选分数最高的 K 个专家
- 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
9GPU
├── 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 编程里有三个重要概念:
- 线程(Thread):最小的执行单元,每个线程处理一小部分数据
- 线程块(Thread Block):一组线程,共享同一个 SM 和共享内存
- 网格(Grid):所有线程块的集合
结构是这样的: 1
2
3
4
5
6
7
8
9Grid(网格)
├── 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 程序的性能,通常受两种因素限制:
- 计算受限(Compute Bound):计算单元不够用,内存带宽还有富余
- 内存受限(Memory Bound):内存带宽不够用,计算单元在闲着等数据
怎么判断是哪种?看算术强度:
1 | |
- 算术强度高 → 计算受限(计算是瓶颈)
- 算术强度低 → 内存受限(内存是瓶颈)
5.3.2 Roofline 图
把这个关系画成图,就是 Roofline 模型:
1 | |
- 左边斜线部分:内存受限,算术强度越低,性能越低
- 右边平顶部分:计算受限,性能达到峰值,不再增长
你的程序的算术强度落在哪个区域,就决定了它的瓶颈是什么。
💡 生活类比: 想象一个工厂: - 工人(计算单元)的生产速度是有限的(计算峰值) - 原材料运输(内存带宽)的速度也是有限的 - 如果产品需要很多原材料(算术强度低),那运输跟不上,工人闲着等(内存受限) - 如果产品需要很少原材料(算术强度高),那工人全力生产,运输不是问题(计算受限)
优化的目标:要么提高算术强度(减少内存访问),要么提高计算效率。
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 的核心思想就是我们刚学的两个技巧:分块 + 算子融合。
具体怎么做?
- 分块:把 Q、K、V 都分成小块
- 逐块计算:每次加载一小块 Q 和一小块 K V 到共享内存
- 在线 softmax:逐块计算 softmax,不需要保存完整的注意力矩阵
- 融合:整个过程在一个 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 的编程模型是三层的:
- Thread(线程):最小的执行单元,每个线程处理一小部分数据
- Thread Block(线程块):一组线程,共享同一个 SM 和共享内存
- Grid(网格):所有线程块的集合
1 | |
为什么要分块?
因为线程之间需要通信(比如计算总和),而通信需要共享内存。同一个块里的线程可以通过共享内存快速通信,不同块之间不能直接通信。
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 | |
这是一个逐元素操作——每个元素独立计算,很适合 GPU 并行。
Triton 代码的结构:
1 | |
解释一下每一步:
- 计算偏移:每个块处理 BLOCK_SIZE 个元素,块编号 pid 决定了处理哪一段
- 掩码(mask):如果元素总数不是 BLOCK_SIZE 的整数倍,最后一个块会越界,用 mask 来屏蔽超出的部分
- 加载:从全局内存读数据到寄存器
- 计算:在寄存器里做计算,这部分很快
- 存储:把结果写回全局内存
怎么调用这个 kernel?
1 | |
[(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 | |
关键点: - 每个块处理一行 - 一整行数据加载到寄存器后,所有计算都在寄存器里完成 - 不需要反复读写全局内存 - 速度比朴素实现快好几倍!
💡 归约操作的特点: - 需要线程间协作(求 max、sum 等) - 如果数据量不大,一整行能放进一个块,就很简单 - 如果数据量大,放不下,就需要分块 + 累积
6.5.3 例子三:行求和(分块归约)
如果一行特别长,一个块放不下怎么办?
这时候就需要分块归约——把一行分成多块,每个块算一部分,最后再合并。
1 | |
思路: 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 | |
解释一下:
- 二维网格:矩阵乘法的网格是二维的——pid_m 是行方向的块编号,pid_n 是列方向的块编号
- 累加器:acc 是 BLOCK_M × BLOCK_N 的矩阵,用来累积结果
- K 维度循环:沿着 K 维度循环,每次加载 A 和 B 的一小块,做矩阵乘法,累加到 acc
- tl.dot:Triton 提供的矩阵乘法指令,会自动用 Tensor Core 加速
- 写回:最后把结果写回全局内存
为什么分块更快? - 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 | |
这样就不需要先把矩阵乘法的结果写回内存,再读出来做 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 | |
越往上越快,但容量越小;越往下越慢,但容量越大。
整个系统的设计原则和单 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 | |
用途:比如初始化时,rank 0 加载模型参数,然后广播给所有 GPU。
💡 生活类比: 老师(rank 0)把讲义发给全班同学(所有 rank),每个人都拿到一份完整的讲义。
2. Scatter(分发)
把 rank 0 的数据分成 N 份,每个 GPU 拿一份。
1 | |
用途:把数据分片,每个 GPU 负责一部分。
💡 生活类比: 班长(rank 0)把一摞作业分成 4 份,每个小组(rank)拿一份去批改。
3. Gather(收集)
Scatter 的反向操作——把所有 GPU 的数据收集到 rank 0。
1 | |
用途:把各个 GPU 的计算结果汇总到一个地方。
💡 生活类比: 各小组批改完作业,把作业交回给班长(rank 0),班长把它们合在一起。
4. Reduce(归约)
所有 GPU 的数据做一个运算(比如求和、求最大值),结果放到 rank 0。
1 | |
运算可以是 sum、min、max 等,只要满足结合律和交换律就行。
用途:比如汇总各个 GPU 的梯度。
💡 生活类比: 每个小组统计自己的捐款金额,班长把所有小组的金额加起来,得到总捐款数。
7.2.2 进阶操作:All-Gather、Reduce-Scatter、All-Reduce
上面四个操作的结果都只在 rank 0,其他 GPU 没有。
但有时候我们希望所有 GPU 都有结果——这就是 "All-" 系列操作。
1. All-Gather(全收集)
Gather 的进阶版——所有 GPU 都收集到完整的数据。
1 | |
相当于每个 GPU 都做了一次 gather。
用途:比如每个 GPU 有一部分参数,需要把所有参数合起来才能做前向传播。
💡 生活类比: 每个小组有一部分拼图,大家把自己的部分传出去,最后每个人都有完整的拼图。
2. Reduce-Scatter(归约-分发)
先 reduce,再 scatter。
每个维度先求和,然后把结果分发到各个 GPU。
1 | |
用途:比如梯度求和后,每个 GPU 只负责更新自己那部分参数。
💡 生活类比: 每个小组有一张成绩表(各科分数),大家先把各科的总分算出来(reduce),然后每个小组负责保管一科的成绩(scatter)。
3. All-Reduce(全归约)
最重要的操作!几乎所有分布式训练都会用到。
All-Reduce = Reduce-Scatter + All-Gather
先 reduce-scatter,再 all-gather,最后每个 GPU 都有完整的归约结果。
1 | |
所有 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 | |
看起来像是矩阵的转置——输入是按行存的,输出是按列存的。
用途:MoE 模型的路由——每个 GPU 有一部分数据和一部分专家,需要把数据送到对应专家的 GPU 上。
💡 生活类比: 想象一个邮局系统: - 每个城市(rank)都有一堆信件,分别寄往不同的城市 - All-to-all 就是把所有信件按目的地分类,每个城市收到寄给自己的所有信件
这是最复杂的通信模式,但也是最灵活的。
7.3 硬件:GPU 之间怎么连?
知道了通信操作,我们来看看硬件层面 GPU 之间是怎么连接的。
7.3.1 节点内:PCIe → NVLink
同一台机器里的 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 分布式训练策略:三种基本方式
有了通信的积木,我们来看怎么用它们来实现分布式训练。
主要有三种基本的并行策略:
- 数据并行(Data Parallelism):切 batch
- 张量并行(Tensor Parallelism):切矩阵的维度
- 流水线并行(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 更复杂,因为参数也是分片的,前向传播时需要先把参数合起来:
- 前向传播:
- 每一层计算前,先 all-gather 这一层的参数(拿到完整参数)
- 计算
- 计算完后,扔掉完整参数(只保留自己的分片)
- 反向传播:
- 每一层计算前,先 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 | |
- 张量并行:同一台机器的 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 | |
每个 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 通用性:专用硬件快但不够灵活
整个系统优化的本质,就是在各种约束条件下找到最优的平衡点。