主要参考Datawhale社区
1 Motivation:常识性问题
1.1 估算时间
在 1024 张 H100 显卡上,训练一个 70B(700亿参数)的模型,数据量是 15T(15万亿 tokens),大概要多久
FLOPs(Floating Point Operations,浮点运算次数)
衡量计算复杂度的常用指标,
模型的权重 和激活值 都是连续实数(比如
0.834、-1.2e-5),而实数在计算机里就是用浮点数 (FP32 / FP16 / BF16)来近似存储的;所以需要用浮点运算次数衡量计算量;该定义表示一个算法或模型在执行过程中所需的浮点加法与乘法运算的总次数。
其中 训练模型本质上是进行浮点运算(FLOPs)。有一个经验公式:
总计算量≈6×参数量×Token数量(前向两倍 反向两倍)
1.1.1 总计算量
总计算量
6×(70× 10 9 10^9 109)×(15× 10 12 10^{12} 1012)≈6.3× 10 24 10^{24} 1024 FLOPs
- B = billion(十亿)= 10⁹
- G = giga(十亿倍)= 10⁹
- 两者大小相同;都是10亿
1.1.2 硬件算力
其 FP16/BF16 的峰值算力约为 1979 TFLOPS (每秒万亿次浮点运算)

1 什么是 2:4 稀疏,为什么能翻倍
从 Ampere 架构(A100) 开始,Tensor Core 支持一种特殊模式:每 4 个连续元素里,最多只有 2 个非零 (所以叫 2:4)。如果权重满足这个结构,硬件就能直接跳过那 2 个为 0 的乘法 ,只算一半的 MAC,于是吞吐量理论上翻倍。
结果就是官方规格表里每档算力都有两个数,稀疏恰好是稠密的 2 倍
2 为什么峰值表总是"报"稀疏那个数
① 营销动机(最主要) :数字越大越好看。稀疏值 = 2× 稠密值,厂商自然愿意把更大的那个当"峰值"印在宣传页上,dense 值反而用小字缩在角落。
② "峰值"本来的定义就是理论极限 :peak 指硬件在理想条件下能达到的最大吞吐。开启 2:4 稀疏后确实能到 2×,所以厂商有底气把它叫作"峰值"------问题是这个理想条件现实里几乎凑不出来。
- 单卡实际算力 ≈ ( 990 × 10 12 ) × 0.5 ≈ 5 × 10 14 (990 \times 10^{12}) \times 0.5 \approx 5 \times 10^{14} (990×1012)×0.5≈5×1014 FLOPS
- 1024 张卡总算力 ≈ 5 × 10 17 5 \times 10^{17} 5×1017 FLOPS
- 时间 = 总工作量 ÷ 总算力 = 6.3 × 10 24 5 × 10 17 ≈ 1.26 × 10 7 \frac{6.3 \times 10^{24}}{5 \times 10^{17}} \approx 1.26 \times 10^{7} 5×10176.3×1024≈1.26×107 秒 ≈ 146 天
1.2 估算显存
在 8 张 H100 GPU 上,使用 AdamW 优化器(朴素实现,即所有数据都用 float32 存储,没有混合精度或压缩),你能训练的最大模型参数量是多少?
- 参数 (FP32): 4 bytes
- 梯度 (FP32): 4 bytes
- 优化器状态 (FP32): 8 bytes (2 个变量 × 4 bytes)
最大参数量 = 640 × 10 9 bytes 16 bytes/param ≈ 400 \dfrac{640 \times 10^{9}\ \text{bytes}}{16\ \text{bytes/param}} \approx 400 16 bytes/param640×109 bytes≈400 亿(40B)
这个计算忽略了激活值内存(取决于 batch size 和序列长度),所以只是一个理论上界 ,实际能训练的模型会比这个小。
激活值 = 前向传播中每一层的输出。数据流过网络时,每经过一个矩阵乘法或非线性函数就产生一批中间结果,这些中间结果就是激活值。
2 张量
2.1 定义
张量(Tensor) 是深度学习和科学计算中最基本的数据结构,可以理解为多维数组的泛化。
- 标量(0维张量) :一个单独的数,如
5 - 向量(1维张量) :一列数,如
[1, 2, 3] - 矩阵(2维张量) :一个二维表格,如
[[1,2],[3,4]] - 张量(3维及以上) :更高维的数组,如形状为
(2,3,4)的三维数组
2.2 操作
2.2.1 视图
理解「视图」是为了在估算显存时分辨哪些操作会真正分配新内存(复制),哪些只是零开销的"重定向"(视图)。
许多操作只是提供张量的一个不同"视图"。这不会创建副本,
两个张量指向同一块内存地址,所以改一个,另一个也会变。
python
x = torch.tensor([[1., 2, 3], [4, 5, 6]])
# 这些操作不复制数据
y = x[0] # 获取第0行
y = x[:, 1] # 获取第1列
y = x.view(3, 2) # 重塑为3x2
y = x.transpose(1, 0) # 转置
# 验证是否共享存储
assert same_storage(x, y) # same_storage检查底层存储是否相同
# 注意:修改x会影响y
x[0][0] = 100
assert y[0][0] == 100 # 值也被修改
注意:并非所有视图都是"连续"(contiguous)的。
- 连续 :张量的元素在底层内存中按行优先顺序紧密排列,逻辑顺序与物理顺序一致(如原矩阵
[[1,2,3],[4,5,6]]存成1,2,3,4,5,6)。- 非连续(Non-contiguous) :
transpose等操作只改变"怎么看"数据(改 stride),底层内存不动,导致按逻辑顺序读取时要"跳步"。例如转置后底层仍是1,2,3,4,5,6,逻辑上却要读成1,4,2,5,3,6。([[1, 4],[2, 5],[3, 6]])
判断依据是 stride :连续要求stride[i] = stride[i+1] × shape[i+1]。
影响 :①view()等要求连续的操作会报错;② 跳步访问内存性能更差;③ 需用.contiguous()复制成连续(有额外内存开销、且不再共享存储)。
转置后的张量是非连续的:
python
x = torch.tensor([[1., 2, 3], [4, 5, 6]])
y = x.transpose(1, 0)
assert not y.is_contiguous() # 非连续张量
你不能对非连续张量直接执行某些操作,比如再次 view:
# 尝试重塑会失败
try:
y.view(2, 3)
except RuntimeError as e:
assert "view size is not compatible with input tensor's size and stride" in str(e)
如果需要对一个非连续张量进行进一步操作,可以先调用 .contiguous() 方法:
# 解决方案:先使张量连续
y = x.transpose(1, 0).contiguous().view(2, 3)
assert not same_storage(x, y) # 现在创建了新存储
.contiguous() 会创建一个新的张量,并将数据按顺序复制到新的连续内存块中,这样后续操作就不会出错。
2.2.2 逐元素操作(Element-wise Operations)
逐元素操作是对每个元素独立应用函数,并返回相同形状的新张量。
与「视图」相反,逐元素操作会产生新值 ,因此会分配新内存 (结果和输入不共享存储)。
主要应用:
- 非线性激活函数:激活函数(ReLU/GELU/sigmoid)都是逐元素操作,是模型能拟合复杂函数的关键------没有它们,多层线性层等价于一层。
- 残差连接 :
x = x + f(x)中的逐元素加法,是 Transformer/ResNet 能训练深层的核心。 - 归一化与掩码:LayerNorm 的减均值/除方差、softmax 的 exp/除法、注意力掩码、dropout,本质都是逐元素运算。
计算特性:
- 完全可并行:每个输出只依赖对应一个输入,完美契合 GPU 的大规模并行设计。
- 通常是 memory-bound :每个元素算术极少(一次乘/加),却要读写整块张量,瓶颈在显存带宽 而非算力------所以推理优化常用算子融合(kernel fusion) ,把
gelu(ln(x + attn(x)))这类一串逐元素操作合并成一次读、一次写,减少带宽压力。
工具函数 triu ,用于 因果注意力掩码(causal attention mask) 。在语言模型中,为了确保模型在预测第 j 个词时只能看到第 j 个词之前的词(即不能"偷看"未来信息),就需要使用这种上三角矩阵作为掩码。其中 Mi, j 表示位置 i 对位置 j 的贡献,当 i > j 时(即 i 在 j 之后),贡献应为0。
# 实用操作:上三角矩阵(用于因果注意力掩码)
x = torch.ones(3, 3).triu() # 创建一个3x3全1矩阵,然后取其上三角部分
# 结果:
# [[1, 1, 1],
# [0, 1, 1],
# [0, 0, 1]]
2.2.3 矩阵乘法(Matrix Multiplication)
矩阵乘法是深度学习的基础。矩阵乘法是神经网络中最核心、最频繁的计算操作,无论是全连接层、卷积层还是注意力机制,其底层都离不开矩阵运算。
一个标准的矩阵乘法形式为:一个 M×K 的矩阵乘以一个 K×N 的矩阵,得到一个 M×N 的结果矩阵。示例如下:
# 基本矩阵乘法
x = torch.ones(16, 32) # 16x32矩阵
w = torch.ones(32, 2) # 32x2权重矩阵
y = x @ w # 结果:16x2矩阵
assert y.size() == torch.Size([16, 2])
2.2.3.1 深入理解:行向量视角下的维度变换
核心图像 :x @ W 就是「用 W 里的 m 个问题,重新询问这个 n 维对象」,把 m 个答案排成新向量。
y j = ∑ i x i ⋅ W i j y_j = \sum_i x_i \cdot W_{ij} yj=∑ixi⋅Wij:每个输出 = 原始数据 x 与 W 某一列(第 j 个"问题")的点积。
W 的每一列是一个"探测器/神经元" ,输出 y j y_j yj = 输入与第 j 列的匹配程度。
全连接层 y = x @ W 中,W 的每一列就是一个神经元的全部权重。
x = [2, 1, 3] (1×3:三个原始特征)
问题0 问题1
W = [ 1 , 0 ] ← 第0列问:"特征0+特征2是多少?"
[ 0 , 1 ] ← 第1列问:"特征1+特征2是多少?"
[ 1 , 1 ]
y = x @ W = [5, 4] (1×2)
变换后旧坐标没了------你活在由 W 的列定义的新坐标上。这就是维度变换的本质:扔掉旧的 n 根坐标轴,换成 W 各列定义的 m 根新轴。
接线图:一台「n 进 m 出」的机器
- 左边缘 n 根进线、右边缘 m 根出线;每根进线连每根出线,连线强度 =
W[i,j] - 出口的值 = 所有进线的加权总和
- 内维相等 = 接口宽度一致才能对接
| 形状 | 名称 | 实例 |
|---|---|---|
| m < n | 压缩/投影 | 词向量 512→2 可视化 |
| m > n | 扩展 | Transformer FFN: 512→2048 再缩回 |
| m = n | 同维换基 | 注意力里的 V 投影 |
batch 维不受影响:(batch, n) @ (n, m) → (batch, m),每个样本独立过机器。
压缩 = 拿信息换效率;扩展 = 拿算力换容量。
选维度本质上是在给这两个东西定价------而深度学习的主流配方是:先扩(给足思考空间)→ 非线性 → 再缩(压缩回工作维度),两头的好处各占一点
2.2.3.2 为什么会消失一个维度?
因为那个维度是求和的方向------对哪个轴求和,哪个轴就被折叠掉。四层理解:
- 点积是最小例子 : 1 , 2 , 3 ⋅ 4 , 5 , 6 = 32 1,2,3 \cdot 4,5,6 = 32 1,2,3⋅4,5,6=32,两个 n 维向量 → 一个数。把 n 个位置的信息汇总成一个数,这根轴就不再是结果的坐标,而是化进了数值里。
- 矩阵乘法 = 批量点积 :结果每个格子都要对整条 K 轴求和( ∑ k A i , k B k , j \sum_k Ai,kBk,j ∑kAi,kBk,j),扫完之后 K 化进格子里,自然不在结果形状中。
- 复合函数的中转站(最深层):B 出口 = A 入口 = K,是两步之间的接口;复合映射 A∘B 直接从入口维走到出口维,中转站只是过程,所以 AB 里它必须消失------否则就不是"一个新变换"了。比如 (2×3,"2 进 3 出"),(3×2,"3 进 2 出)
- 规律:消失的永远是被求和的轴。逐元素乘不求和 → 形状不变;点积/矩阵乘/sum 都内置求和 → 被求和的轴消失。
术语 :对重复出现的轴求和并消掉叫缩并(contraction) 。
einsum 口诀:哪个下标在输入中出现两次、又没写进箭头右边,哪个维度就消失。
2.2.3.3 两种书写约定对照
| 教科书(列向量) | PyTorch(行向量) | |
|---|---|---|
| 写法 | y = A x y = Ax y=Ax | y = x @ W |
| 单个输出元素 | x 点乘 A 的行 | x 点乘 W 的列 |
| 探测单元(神经元) | A 的行 | W 的列 |
| 复合阅读顺序 | A B x ABx ABx 从右往左读(先 B 后 A) | x @ A @ B 从左往右读 |
两套约定互为转置 ( A x ) T = x T A T (Ax)^T = x^T A^T (Ax)T=xTAT,本质同一件事。唯一不变的硬规则:内维必须相等------它是被求和消掉的接口维度。
2.3 Einops 库
2.3.1 问题
在 pytorch 中,张量维度通常是 [batch, sequence, hidden]。用 PyTorch 原生的 .view() 和 .transpose() 操作维度时,需要时刻记住张量的维度顺序,很容易搞晕。并且,如果你未来修改了张量形状(比如在 transformers 模型中加了 heads 维度),这个代码就可能出错。
x = torch.ones(2, 2, 3) # batch, sequence, hidden
y = torch.ones(2, 2, 3) # batch, sequence, hidden
z = x @ y.transpose(-2, -1) # 得到 (batch, sequence, sequence)
几个工具(使用 jaxtyping 先声明维度含义,再使用 einops 来操作张量),让代码像写公式一样清晰。
2.3.2 用 jaxtyping 命名维度
# 传统方式(容易写错维度顺序)
x = torch.ones(2, 2, 1, 3) # batch seq heads hidden
# Jaxtyping方式(在类型注解中命名维度)
from jaxtyping import Float
x: Float[torch.Tensor, "batch seq heads hidden"] = torch.ones(2, 2, 1, 3)
jaxtyping 提供的维度命名(如 "batch seq hidden")在当前 PyTorch 生态中主要是"文档性质"的,不会在运行时自动强制检查维度是否真的匹配名称,
也就是说PyTorch 在运行时不会检查 x 的形状是否真的是 (batch=2, seq=3, hidden=4),即使你写成:
y: Float[torch.Tensor, "batch seq hidden"] = torch.randn(100, 5) # 形状与注解严重不符
代码依然会正常运行。
Python 的类型注解(Type Hints)本身是可选的,主要用于工具(如 IDE、mypy)做静态分析。
jaxtyping能极大地提升代码的清晰度和可靠性。
当将来修改模型(比如增加 heads 维度):
- 类型注解会提醒你哪里需要同步更新。
- IDE(如 VS Code、PyCharm)能基于注解提供:
- 自动补全
- 重构支持(rename dimension)
- 错误高亮(如果你在 einsum 中拼错了维度名)
注意:
jaxtyping也提供了可选地运行时检验,第三方库如 beartype 也提供了类型注解做运行时校验的功能
2.3.3 einops.einsum 替代矩阵乘法 + 转置
广义矩阵乘法,具有清晰的维度变化。
这本质上是爱因斯坦求和约定(Einstein summation) 的直观实现。
什么是爱因斯坦求和约定? 爱因斯坦为省去连加号 ∑ 提出的书写约定:同一个下标在乘积项中出现两次,就自动对该下标求和 。如矩阵乘法 C i k = ∑ j A i j B j k C_{ik}=\sum_j A_{ij}B_{jk} Cik=∑jAijBjk 简写为 C i k = A i j B j k C_{ik}=A_{ij}B_{jk} Cik=AijBjk(
j重复 → 隐式求和)。einsum 就是它的直观实现 :把该约定翻译成字符串语法------
"ij, jk -> ik"中j重复表示求和,i/k表示保留的输出维度。einops.einsum进一步用命名维度("hidden")取代无意义的i/j/k。本质 :一种「声明式」的张量运算描述------只需声明哪些轴相乘、哪些轴缩并、结果保留哪些轴 ,无需关心具体形状或是否要先转置,天然通用、支持广播与前导维度(
...),不易写错维度顺序。
更通用的写法是(支持任意前导维度):
z = einsum(x, y, "... seq1 hidden, ... seq2 hidden -> ... seq1 seq2")
... 表示"任意数量的前导维度",比如可能是 (device, batch) 或 (ensemble, batch, time),代码依然适用。好处是维度逻辑显式、通用、不易出错,且天然支持广播。(:两个输入的前导维度(... 那部分)即使形状不完全一样,einsum 也会按名字自动对齐、缺的维度自动广播 ,不用你手动去 .unsqueeze() 或 .expand() 补维度。)
from einops import einsum
# 定义两个张量
x: Float[torch.Tensor, "batch seq1 hidden"] = torch.ones(2, 3, 4)
y: Float[torch.Tensor, "batch seq2 hidden"] = torch.ones(2, 3, 4)
# 传统方式
z = x @ y.transpose(-2, -1) # batch, sequence, sequence
# Einops方式
z = einsum(x, y, "batch seq1 hidden, batch seq2 hidden -> batch seq1 seq2")
# 未在输出中命名的维度(hidden)会被自动求和
2.3.4 用 einops.reduce 替代 mean(dim=...)
前面 einsum 解决的是"相乘 + 求和",但很多归约操作不止求和------比如对某个维度求平均 。传统写法 x.mean(dim=-1) 用的是位置索引,既不直观(-1 是哪一维?)又容易在维度顺序变化时静默出错。einops.reduce 用和 einsum 一致的命名维度语法,把"对哪个维度、做什么归约"写清楚。
from einops import reduce
x: Float[torch.Tensor, "batch seq hidden"] = torch.ones(2, 3, 4)
# 传统写法
y = x.mean(dim=-1) # 对最后一个维度 hidden 上求平均
# Einops 写法
y = reduce(x, "... hidden -> ...", "mean")
从 ... hidden 变成 ...,说明 hidden 维度被"聚合"了,聚合方式是 "mean"(也可用 "sum", "max" 等)。 同样... 表示"任意数量的前导维度"。这种写法的好处是明确表达了"我打算把哪个维度压缩掉",语义清晰。
2.3.5 用 einops.rearrange 拆分/合并维度
拆分/合并维度(如多头注意力里把 heads×dim 拆成 heads 和 dim)是深度学习里最容易写错的地方:原生 view + transpose 步骤多、顺序易错,且 transpose 后张量非连续,再 view 会报错、还得手动 .contiguous()。rearrange 用一行命名维度的声明式语法同时搞定拆分、交换、合并,底层自动处理转置与连续性。
假设你有一个维度 total_hidden = 8,它实际上是 heads=2 和 hidden1=4 的乘积(即 2 * 4 = 8),也就是 total_hidden 维度实际上是 heads×hidden1 的扁平化表示。
from einops import rearrange
# 情景:hidden维度实际上是heads×hidden1的扁平化表示
x: Float[torch.Tensor, "batch seq total_hidden"] = torch.ones(2, 3, 8)
w: Float[torch.Tensor, "hidden1 hidden2"] = torch.ones(4, 4)
为什么要这样设计? 单头注意力只有一组 Q/K/V,所有 token 之间的相关性只能按一种模式衡量;但语言里的"相关"是多样的------语法关系、指代关系、局部邻近等,一组投影学不下。所以把 d_model 拆成 heads × head_dim,让每个头拥有一组独立的 Q/K/V,在各自的子空间里并行学习不同的注意力模式(类比 CNN 里不同卷积核提取不同特征)。
拆头之后,每个头做一次矩阵乘法 = 用各自的 Q/K 算"谁关注谁"(Q·Kᵀ),再用 V 加权聚合,独立算一遍注意力;并头 = 把各头输出拼接回 d_model,合并成一个完整表示交给下一层。
对应 Transformer 步骤 :下面三步正是多头注意力(Multi-Head Attention)里「拆头 → 每头计算 → 并头」的完整流程。注意力中 d_model = heads × head_dim 是扁平存放的,计算前要拆成独立的多头维度,算完再拼回去。
① 拆头(split) :(batch, seq, heads×head_dim) → (batch, seq, heads, head_dim)
(heads hidden1) 表示「这两个维度原本被乘在一起,现在拆开」。因为 8 能拆成 (2,4)/(4,2)/(8,1) 等,必须用 heads=2 固定拆分方式。
python
x = rearrange(x, "... (heads hidden1) -> ... heads hidden1", heads=2)
② 每个头做矩阵乘法 :拆完后 heads 作为前导维度,einsum 对 hidden1 求和,相当于每个头各自做一次变换:
python
x = einsum(x, w, "... hidden1, hidden1 hidden2 -> ... hidden2")
③ 并头(merge) :(batch, seq, heads, head_dim) → (batch, seq, heads×head_dim),反向括号即合并:
python
x = rearrange(x, "... heads hidden2 -> ... (heads hidden2)")
拆分与合并正好是
( )括号的两个方向:括号内 = 拆开,括号外 = 合并 ,与多头注意力的拆头/并头一一对应。补充:真实 Transformer 里拆头后还要把 heads 提到前面(转置),
rearrange可一步完成「拆头 + 转置」:"batch seq (heads dim) -> batch heads seq dim",这也是它相比view + transpose最省心的地方。
尽管Einops增加了少量语法开销,但其清晰的维度命名显著降低了调试难度,特别是在复杂的模型架构中,更多用法可见
3 内存(Memory)
3.1 因素
张量的内存占用由两个因素决定:
- 元素数量:张量的形状(如4×8矩阵有32个元素)
- 数据类型:每个元素占用的字节数
上面第二个因素「数据类型」具体指的就是浮点类型------不同浮点类型每个元素占用的字节数不同(FP32 占 4 字节、FP16/BF16 占 2 字节、FP8 占 1 字节),因此同样元素数量下,选不同浮点类型,内存占用会相差数倍。下面逐一介绍常见浮点类型。
3.2 浮点类型
3.2.1 为什么浮点数分「符号位 / 指数位 / 尾数位」?
浮点数要同时表示极大、极小、负数、小数,靠的是科学计数法 :任何数都能写成 符号 × 尾数 × 基数^指数(如 -1.2345 × 10³)。
三部分各司其职:
- 符号位:正负(0 正、1 负)
- 指数位:决定数量级 / 小数点位置("这数多大")------每多 1 位,可表示范围翻倍
- 尾数位 :决定有效数字("这数多精确")------每多 1 位,精度翻倍
三者如何拼成一个值(IEEE 754 规范化数公式):
值 = ( − 1 ) 符号 × ( 1. 尾数 ) × 2 指数 − 偏置 \text{值} = (-1)^{\text{符号}} \times (1.\text{尾数}) \times 2^{\text{指数} - \text{偏置}} 值=(−1)符号×(1.尾数)×2指数−偏置
- 尾数隐含前导
1(规范化后首位恒为 1,省去不存,白赚 1 bit 精度); - 指数存「真实指数 + 偏置」(FP32 偏置 = 127),避免单独存指数的正负号。
例:FP32 表示 6.5 : 6.5 = 1.101 2 × 2 2 6.5 = 1.101_2 \times 2^2 6.5=1.1012×22 →
下标 _2 意思是「这是二进制数 」,不是小数 1.101。写成十进制它其实是:
(二进制小数点后第一位是 ½、第二位是 ¼、第三位是 ⅛,所以 1.101 = 1×1 + 1×½ + 0×¼ + 1×⅛ = 1.625)
是 1.625 × 4 = 6.5**。
符号 0、指数存 2+127=129、尾数 101,拼成 0 10000001 10100000000000000000000,还原得 1.101 2 × 2 2 = 6.5 1.101_2 \times 2^2 = 6.5 1.1012×22=6.5 ✓
3.2.2 Float32 (FP32 / 单精度浮点数)

- 规格 :占用 4 字节 (32 bits)。结构为:1位符号位 + 8位指数位 + 23位尾数位。
- 地位:PyTorch 的默认数据类型,也是科学计算领域的"黄金标准"。
- 优点:数值精度高,动态范围大,训练最稳定,几乎不会出现数值溢出问题。
- 缺点:对于大模型而言太"奢侈"。它占用的显存是 16位格式的两倍,且在现代 GPU(如 H100)上的计算吞吐量远低于低精度格式。
- 炼丹用途 :通常用于存储参数的主副本 (Master Weights) 和 优化器状态 (Optimizer States),以确保在梯度累积和参数更新时不会因为精度丢失而导致模型无法收敛。
3.2.3 Float16 (FP16 / 半精度浮点数)

- 规格 :占用 2 字节 (16 bits)。结构为:1位符号位 + 5位指数位 + 10位尾数位。
- 优点:显存占用比 FP32 减少一半,计算速度快。
- 致命缺陷 :动态范围(Dynamic Range)太窄 。
- 由于指数位只有 5 位,它无法表示非常小的数(会发生下溢 Underflow,直接变成 0)或非常大的数(会发生上溢 Overflow,变成 Infinity)。
- 例如:在 FP16 中,
1e-8这样的小数会被直接当作0处理,导致梯度消失。
- 用途 :这是上一代 GPU(如 V100)混合精度训练的主流。为了解决溢出问题,必须使用复杂的损失缩放 (Loss Scaling) 技术。目前在 LLM 训练中正逐渐被 BF16 取代。
3.2.4 BFloat16 (BF16 / Brain Floating Point)

- 规格 :占用 2 字节 (16 bits)。结构为:1位符号位 + 8位指数位 + 7位尾数位。
- 来源:由 Google Brain 专为深度学习设计。
- 设计逻辑 :"要范围,不要精度" 。
- 深度学习模型(尤其是神经网络)对小数点的后几位精度不敏感,但对数值的范围非常敏感。
- BF16 直接截断了 FP32 的尾数,但保留了和 FP32 相同的 8 位指数位。
- 优点 :
- 拥有和 FP32 一样宽广的动态范围,不需要 Loss Scaling 也能稳定训练。
- 显存占用和 FP16 一样少。
- 在 A100/H100 等新硬件上计算速度极快。
- 用途 :当前 LLM 训练的绝对主流选择 。通常用于存储激活值 (Activations) 以及进行前向和反向传播的矩阵乘法计算。
Percy 在课上解释说:"bf16 牺牲了精度来换取范围。对于深度学习来说,范围比精度重要得多,因为数值稳定性的主要威胁是溢出/下溢,而不是尾数精度不够。"
最大值 ≈ 2 2 E − 1 2^{\,2^{E-1}} 22E−1 ( E E E = 指数位数)
推导思路(简化版):
- 尾数部分最多约等于 2(
1.111...→ 接近 2) - 指数的最大有效值 ≈ 2 E − 1 2^{E-1} 2E−1(因为有偏置,且全 1 要保留给 inf/NaN)
- 所以 max ≈ 2 × 2 2 E − 1 ≈ 2 2 E − 1 2 \times 2^{2^{E-1}} \approx 2^{2^{E-1}} 2×22E−1≈22E−1
- E = 5 E=5 E=5(FP16)→ 2 16 2^{16} 216 ≈ 6 万
- E = 8 E=8 E=8(BF16/FP32)→ 2 128 2^{128} 2128 ≈ 10 38 10^{38} 1038
3.2.5 FP8 (8位浮点数)
-
规格 :占用 1 字节 (8 bits)。
-
变体:
- E4M3:4位指数,3位尾数(精度稍高,范围稍小)。
- E5M2:5位指数,2位尾数(范围稍大,精度更低)。
-
优点:极致的显存压缩和计算吞吐量。
-
硬件限制:仅在 NVIDIA H100 及更新的架构(配合 Transformer Engine)上才被原生支持。
-
炼丹用途 :目前主要用于推理阶段的量化 (Quantization) 。虽然理论上可以用于训练(H100 支持),但由于精度极低,训练极不稳定,属于比较前沿的研究领域。

-
规格 :占用 1 字节 (8 bits)。
-
变体:
- E4M3:4位指数,3位尾数(精度稍高,范围稍小)。
- E5M2:5位指数,2位尾数(范围稍大,精度更低)。
-
优点:极致的显存压缩和计算吞吐量。
-
硬件限制:仅在 NVIDIA H100 及更新的架构(配合 Transformer Engine)上才被原生支持。
-
炼丹用途 :目前主要用于推理阶段的量化 (Quantization)。虽然理论上可以用于训练(H100 支持),但由于精度极低,训练极不稳定,属于比较前沿的研究领域。
3.2.6 FP4 / NVFP4(4 位浮点)
2025 年,NVIDIA 开发了 NVFP4,每个值仅 4 bits!
3.2.7 NVFP4 拆解(导读)
3.2.7.1 为什么只有 16 个值?
4 个 bit 只能组合出 2 4 = 16 2^4 = 16 24=16 种状态(0000~1111),所以只能表示 16 个数------这是硬限制。
具体是哪 16 个,由 E2M1 格式决定(1 符号 + 2 指数 + 1 尾数):
| 指数位(2bit) | 尾数(1bit) | 正数值 |
|---|---|---|
00 |
0 / 1 | 0 / 0.5 |
01 |
0 / 1 | 1.0 / 1.5 |
10 |
0 / 1 | 2.0 / 3.0 |
11 |
0 / 1 | 4.0 / 6.0 |
正的 8 个:0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0,加负号共 16 个。
关键 :这些值不均匀 ------越靠近 0 越密、越远离越稀(0.5→1→1.5→2 间隔 0.5,3→4→6 间隔 1、2)。这是浮点数"指数管范围、尾数管精度"在 4 bit 下的极端体现。
3.2.7.2 训练中如何用?------Block 量化
若真只用这 16 个固定值,权重里大量 0.001 这种小值会全部被舍入成 0,无法训练。所以实际是分块量化 :
把权重切成小块(每 16~32 个值一块),每块共享一个高精度缩放因子 s。流程四步:
- 分块:权重矩阵切成若干小块;
- 算 s :取块内最大绝对值,如
s = 0.004(用高精度单独存); - 归一化 + 舍入:块内每个值除以 s,再舍入到 16 个档位;
- 反量化:真实值 = 存的 4-bit 值 × s。
举例 :权重 [0.001, 0.002, 0.003, 0.004],不用 block 时全变 0;用 block(s=0.004)后归一化为 [0.25, 0.5, 0.75, 1.0],存 [0.5, 0.5, 0.5, 1.0],恢复为 [0.002, 0.002, 0.002, 0.004]------有误差,但量级和结构保住了 。
为什么范围扩大了:4-bit 值管"块内相对大小",s 管"整块量级"。s 可以是任何数,所以不同块可以处在相差千万倍的量级上,绕开了 4 bit 范围窄的限制。
类比:4 bit = 一把只有 16 个刻度的尺子。block 化就是每块换一把单位不同的尺子 (这块用毫米、那块用公里)。
代价:同一块共享一个 s,所以不能"一个值超大、旁边超小"------它们被迫用同一把尺子。这就是为什么块要切得小、且常按通道/分组切(同通道权重往往同量级)。
3.2.7.3 缩放因子 s 用什么格式存?
E8M0 :8 位指数、0 位尾数------s 永远是 2 的幂。
- 因为乘 2 k 2^k 2k 在二进制里就是"指数加 k",硬件上几乎免费(无需真乘法器);
- s 本来只管量级(范围),不需要尾数(精度),所以只存指数就够了------这正好是"指数管范围"的极端化应用。
| 方案 | 数据格式 | s 的格式 |
|---|---|---|
| NVFP4 / MXFP4 | 4-bit E2M1 | 8-bit E8M0(2 的幂) |
| int4/int8(AWQ/GPTQ) | 4/8-bit 整数 | FP16 / FP32 |
| MXFP8 | 8-bit E4M3/E5M2 | 8-bit E8M0 |
通用规律:s 的精度一定 ≥ 数据精度------数据粗一点无所谓,但 s 决定整块量级,不能丢。
2026 年发布的 Nemotron 3 Super 是在 NVFP4 精度下训练的大型模型。
Percy 也提到,fp4 这些底层精度操作实际上是在 NVIDIA 的软件栈中自动完成的,"不是你创建一个张量然后调用 tensor.fp4() ------ 很多工作是在底层 under the hood 进行的,用户无法直接控制。"
不同精度的运算速度完全不同。Percy 特别强调,现在的 GPU 已经不太优化 fp32 了:"如果你现在用 fp32 做训练,会发现真的非常非常慢,因为硬件优化的重点已经转向了 bf16 甚至 fp8。"
3.2.8 各代 GPU 支持的浮点格式一览
每一代新架构都会新增一种更低的精度格式------这是理解 GPU 演进最清晰的线索:
| 架构 | 代表卡 | 发布时间 | FP32 | TF32 | FP16 | BF16 | FP8 | FP4 |
|---|---|---|---|---|---|---|---|---|
| Volta | V100 | 2017.05 | ✅ | ❌ | ✅ | ❌ | ❌ | ❌ |
| Turing | T4 / RTX 20 系 | 2018.09 | ✅ | ❌ | ✅ | ❌ | ❌ | ❌ |
| Ampere | A100 / RTX 30 系 | 2020.05 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| Ada Lovelace | L40S / RTX 40 系 | 2022.10 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ |
| Hopper | H100 / H200 | 2022.03(H200:2024) | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ |
| Blackwell | B200 / GB200 / RTX 50 系 | 2024.03(出货 2024 底~2025) | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
后续:Rubin 架构预计 2026 年推出,延续低精度路线。
规律与要点:
- 每约 2 年一种新低精度格式:FP16(Volta) → TF32/BF16(Ampere) → FP8(Hopper) → FP4/NVFP4(Blackwell),精度减半、吞吐翻倍;
- TF32 与 BF16 都是 Ampere 起才有硬件支持------V100 上跑不了原生 BF16;
- FP8 从 Hopper 开始(配合 Transformer Engine 自动管理),消费级的 Ada 也支持;
- FP4 是 Blackwell 独有(NVFP4/MXFP4),目前主要用于推理量化,训练刚起步(如 Nemotron 3 Super);
- 所以选精度前先看卡:老卡(V100/T4)上谈 FP8/FP4 训练是没有意义的。
3.3 张量在内存中的存储机制
在 PyTorch 的内部实现中(继承自 Lua Torch 的设计,直到 1.x 仍是公开 API):
张量占用了两部分内存------头信息区(Tensor) 储存形状(size)、步长(stride)、数据类型(dtype)等元数据;存储区(Storage) 储存真正的数据。
但将 Storage 作为可公开访问的对象,其带来的复杂性和潜在风险远大于灵活性,因此 从 PyTorch 2.0 起已基本将其视为内部实现细节 ,用户不应该再直接操作 .storage()。理解它主要是为了看懂底层和多张量共享内存的原理。
PyTorch 的设计者们选择将这一层实现细节进行封装,让 Tensor 对象本身成为一个集成了"指针+元数据"的单一实体。这意味着一个张量本身并不直接"包含"数据,而是指向一块连续的内存区域,并附带一套规则(元数据),告诉程序如何根据你请求的索引去这块内存里找到对应的数据。
markdown
- **早期设计**:暴露两个对象------张量(元数据:形状/dtype)+ 存储区(裸数据)。用户可直接操作存储区。
- **现在的设计**:张量 = 指针 + 元数据,打包成**单一实体**;存储区降级为内部实现细节,用户摸不到。
让我们定义一个4x4的二维张量作为例子
x = torch.tensor([
[0., 1, 2, 3],
[4, 5, 6, 7],
[8, 9, 10, 11],
[12, 13, 14, 15],
])
张量在底层可以是按行存储也可以是按列存储。Numpy 和 Pytorch 都采用了按行存储的方式,任何维度的张量在底层存储都占据着内存中连续的空间
3.4 张量在cpu与GPU内存移动

上图左侧包含中央处理器(CPU)和系统内存(RAM),右侧包含图形处理器(GPU)及其专用的高速显存(DRAM)
GPU内部由多个"流式多处理器"(streaming multiprocessor)组成,这是其并行计算能力的来源。
CPU和GPU通过"PCI BUS"(PCI总线)相连,数据在CPU和GPU之间传输需要经过这条总线,这通常是一个性能瓶颈,因为它的带宽远小于GPU内部的内存带宽。因此,在实际应用中,我们应尽量减少CPU和GPU之间的数据传输
python
# ============ 将张量从 CPU 移至 GPU ============
import torch
# 首先,检查当前机器上有没有 GPU
if not torch.cuda.is_available():
return # 没有 GPU 则直接返回(注意:该写法仅适用于函数内部)
# ---- 获取 GPU 信息 ----
num_gpus = torch.cuda.device_count() # 获取 GPU 数量
for i in range(num_gpus):
# 获取第 i 块 GPU 的属性(如型号、显存大小等)
properties = torch.cuda.get_device_properties(i)
# 记录当前已分配给 PyTorch 的显存总量(作为基准)
memory_allocated = torch.cuda.memory_allocated()
# ---- 移动张量到 GPU 有两种方法 ----
# 方法一:移动现有张量
y = x.to("cuda:0") # 将张量 x 移动到编号为 0 的 GPU 上
assert y.device == torch.device("cuda", 0) # 验证移动成功
# 方法二:直接在 GPU 上创建张量
z = torch.zeros(32, 32, device="cuda:0") # 在 GPU 上直接创建一个 32x32 的零矩阵
# 这种方法更高效:避免了"先在 CPU 创建再搬运"的过程
# ---- 最后,验证内存分配 ----
new_memory_allocated = torch.cuda.memory_allocated() # 再次查询当前显存占用
memory_used = new_memory_allocated - memory_allocated # 计算新增的显存占用
assert memory_used == 2 * (32 * 32 * 4) # 新增了 2 个 32x32 矩阵,每个元素是 4 字节的 float
4 计算效率
4.1 浮点运算总数(FLOPs)
4.1.1 概述
- FLOPs(小写s):执行的浮点运算总数,这是一个衡量计算量的指标,
- FLOP/s (每秒浮点运算):指硬件每秒能执行的浮点运算次数,也写作 FLOPS(大写S)。这是一个衡量"速度"的指标。例如,H100 GPU峰值性能为 990 TFLOP/s。
重要提示:这两个缩写写法相似,但含义完全不同!本教程中会明确区分它们,避免混淆。
通过几个例子让你对 FLOPs 和 FLOP/s 有一个直观认识:
4.1.1.1 代表性模型的训练计算量
| 模型 | 年份 | 参数量 | 训练数据 | 训练总计算量 | 来源 |
|---|---|---|---|---|---|
| GPT-3 | 2020 | 175 B | ~3000 亿 tokens | 3.14 × 10 23 3.14 \times 10^{23} 3.14×1023 FLOPs | Lambda 文章 |
| PaLM | 2022 | 540 B | ~7800 亿 tokens | 2.56 × 10 24 2.56 \times 10^{24} 2.56×1024 FLOPs | 官方论文 |
| GPT-4 | 2023 | 未公开(推测 ~1.8T MoE) | ~13 万亿 tokens(推测) | ≈ 2 × 10 25 \approx 2 \times 10^{25} ≈2×1025 FLOPs | 推测文章 |
| DeepSeek-V3 | 2024 | 671 B MoE(激活 37 B) | 14.8 万亿 tokens | 278.8 万 H800 时 ≈ 3.3 × 10 24 3.3 \times 10^{24} 3.3×1024 FLOPs | 官方报告 |
| Llama 3.1 405B | 2024 | 405 B | 15.6 万亿 tokens | 3.8 × 10 25 3.8 \times 10^{25} 3.8×1025 FLOPs | Meta 官方 |
四年间训练计算量涨了 ~100 倍------这就是「规模定律」驱动下算力需求的基本盘。
4.1.1.2 代表性 GPU 的峰值算力
| GPU | 架构 / 年份 | FP16/BF16 稠密 | FP16 稀疏 | FP8 稠密 | 新增能力 |
|---|---|---|---|---|---|
| V100 | Volta / 2017 | 125 TFLOPS | ❌ | ❌ | 首代 Tensor Core |
| T4 | Turing / 2018 | 65 TFLOPS | 130 TFLOPS | ❌ | INT8/INT4 量化推理 |
| A100 | Ampere / 2020 | 312 TFLOPS | 624 TFLOPS | ❌ | TF32、BF16、结构化稀疏 |
| H100 | Hopper / 2022 | 989.5 TFLOPS | 1979 TFLOPS | 1979 TFLOPS | FP8 + Transformer Engine |
| B200 | Blackwell / 2024 | ~2250 TFLOPS | ~4500 TFLOPS | ~4500 TFLOPS | FP4/NVFP4(稠密 ~9 PFLOPS) |
| H20 | Hopper 特供 / 2024 | 148 TFLOPS | 296 TFLOPS | 296 TFLOPS | 算力阉割、显存保留 |
来源:A100 手册 · Tensor Core 手册
注意 :H100 的 1979 TFLOPS 是开启稀疏 后的数字;实际训练中的稠密矩阵乘法约为其一半(989.5 TFLOPS)。估算时间时永远用稠密值再乘利用率。
H20(中国特供版)的启示 :受出口管制限制,其稠密 FP16 算力只有 148 TFLOPS------约为 H100 的 15% ;但显存反而更大(96 GB HBM3,带宽 4.0 TB/s)。这个「砍算力、留带宽」的配置组合,正好印证了前面学的瓶颈分类:训练是 compute-bound (所以训练性能大打折扣),推理是 memory-bound(带宽和显存才是关键,因此 H20 主要用于推理场景)。一张卡的规格表就是一道「瓶颈分析」考题。
4.1.2 线性模型的计算量
既定模型有 B 个数据点,每个点是 D 维的,模型将其映射到 K 维输出。我们要做的核心操作是矩阵乘法 y = x @ w,目标是计算这个操作总共需要多少次浮点运算(FLOPs)?
if torch.cuda.is_available():
B = 16384 # Number of points
D = 32768 # Dimension
K = 8192 # Number of outputs
else:
B = 1024
D = 256
K = 64
device = get_device()
x = torch.ones(B, D, device=device)
w = torch.randn(D, K, device=device)
y = x @ w
考虑 y = x @ w 中的一个输出元素 yi, k:
y[i, k] = x[i, 0] * w[0, k] +
x[i, 1] * w[1, k] +
...
x[i, D-1] * w[D-1, k]
这个求和过程包含 D 次乘法(xi, j * wj, k)和 D - 1 次加法,总共 ≈ 2D 次 FLOPs(因为 D 通常很大,D - 1 ≈ D)
因此,矩阵乘法的计算量:总 FLOPs ≈2×B×D×K
因为 D × K 正好是这个线性层的参数数量!所以我们可以重写为: F L O P s = 2 × 数据量 × 模型参数量 FLOPs = 2 \times 数据量 \times 模型参数量 FLOPs=2×数据量×模型参数量
或者在语言模型中,常用 token 代替数据点: F L O P s = 2 × t o k e n 数量 × 模型参数量 FLOPs = 2 \times token 数量 \times 模型参数量 FLOPs=2×token数量×模型参数量。
虽然相较于线性模型, Transformer 还有注意力、LayerNorm 等额外操作,但矩阵乘法占主导,所以这个公式对 Transformer 也基本成立
4.1.3 其他操作的 FLOPs
- 逐元素操作(Element-wise Operations):对张量中每个元素独立应用某个函数(如 x + 1、x.sqrt()、torch.relu(x) 等)。若张量形状为 m × n,则逐元素操作的 FLOPs 为 O(m × n)。
- 张量加法(Tensor Addition):两个相同形状张量的对应元素相加,如 z = x + y。若 x 和 y 均为 m × n,则总 FLOPs = m × n(每个元素 1 次加法)。
对于现代深度学习模型(尤其是 Transformer、MLP 等),90%+ 的 FLOPs 来自矩阵乘法(线性层、QKV 投影、FFN 等)。其他操作(LayerNorm、Softmax、激活函数、加法残差连接等)虽然存在,但 FLOPs 可忽略不计。因此,在进行粗略但有效的资源估算时,只需计算所有矩阵乘法的 FLOPs 总和即可 。
将 FLOPs 转化为实际运行时间
actual_time = time_matmul(x, w) # 实际执行矩阵乘法所需的时间(秒)
actual_flop_per_sec = actual_num_flops / actual_time # 实际每秒能完成多少次浮点运算
FLOP/s 的性能很大程度上取决于数据类型!
promised_flop_per_sec = get_promised_flop_per_sec(device, x.dtype) # 获取硬件的理论峰值性能
为什么数据类型影响这么大
① 跑的计算单元不一样
- FP32 走通用 CUDA Core------它要兼顾各种任务,矩阵乘法不是它的主场;
- FP16/BF16/FP8 走 Tensor Core ------专门为矩阵乘加设计的电路,一个周期算一整个小矩阵。
这就是 Percy 那句"现在 GPU 已经不太优化 fp32 了"的硬件根源:厂商把晶体管都堆到 Tensor Core 上了。
② 精度越低,同样面积塞得越多
一个 FP64 乘法器需要的电路面积,大约是 FP32 的好几倍、FP8 的几十倍。位数砍半 → 加法器/乘法器面积骤减 → 同样的硅片面积上能并行摆更多计算单元 → 吞吐翻倍。
这正好呼应你笔记里那张表:每约两年新增一种更低精度格式,吞吐随之翻倍 (FP16 → FP8 → FP4),本质就是"用更少的位换更多的并行度"。
③ 低精度还省带宽
每个数从 4 字节变 1 字节,搬运量降为 1/4------计算单元更容易"喂饱"(还记得 decode 阶段是带宽瓶颈吗?低精度直接缓解这个)。
4.1.4 实际计时与 Benchmarking
要知道实际花了多长时间,需要实际测量。Percy 介绍了 benchmark 的基本要素:
python
def benchmark(func, num_trials=5):
if torch.cuda.is_available():
torch.cuda.synchronize() # 确保之前的 CUDA 操作已完成
# 运行 num_trials 次并取平均
...
torch.cuda.synchronize() # 操作后的同步点
关键注意事项:
torch.cuda.synchronize()是必须的:GPU 操作是异步的,不加同步点你会发现"哇,好快"------其实只是操作还没开始执行,调用就返回了。- 通常需要多次运行取平均值来减少噪声。
同步 vs 异步执行
一句话定义:
同步执行(synchronous) :发出任务后原地等待,任务彻底完成才返回,继续执行下一行------"不干完不走人"。
异步执行(asynchronous) :发出任务后立刻返回 ,任务在后台进行,之后再来"收割"结果------"先下单,菜好了再叫号"。
奶茶店类比 :同步 = 站在柜台前盯着店员做完才走;异步 = 扫码下单拿个号就走,做好了叫号再来取。下单那一刻两个世界就分叉了。
代码视角:同步:总耗时 = 任务1 + 任务2 + 任务3(串行累加)
r1 = task1() # 等它做完
r2 = task2()
r3 = task3()异步:三行瞬间返回,任务后台并行
f1, f2, f3 = submit(task1), submit(task2), submit(task3)
result = wait(f1) + wait(f2) + wait(f3) # 需要结果时才等异步总耗时 ≈ 最慢的那个任务,而不是三者之和------这是性能优势的来源。
对应 GPU 场景:
y = x @ w是异步的:CPU 把 kernel 扔进队列后立刻返回,GPU 慢慢算;
torch.cuda.synchronize()就是手动插入的同步点 ------把"异步的世界"摁回"同步的时刻"。
为什么 GPU 要选异步 :CPU 发起一次 kernel 只要几微秒,GPU 执行可能要几毫秒------如果每次都同步等,CPU 和 GPU 会互相干瞪眼,谁也吃不饱。
一图总结:同步: 发起 ────等待────▶ 完成 ──▶ 下一件事
异步: 发起 ──▶ 返回(立刻)──▶ 下一件事照常跑
└──后台:任务执行中......──▶ 完成 ──▶ 收割结果记忆:同步 = 当面等结果;异步 = 先拿小票,稍后再取货。
为什么必须 synchronize?------ GPU 是异步执行的
在 PyTorch 里执行一行 CUDA 操作(如 y = x @ w),CPU 只做两件事:把 kernel 扔进 GPU 的任务队列,然后立刻返回 继续跑下一行 Python------根本不等 GPU 算完。这是故意设计:CPU(服务员)负责下单,GPU(后厨)炒菜,CPU 源源不断投喂才能让 GPU 不闲着。
但异步会让计时彻底失真:
python
start = time.time()
y = x_big @ w # GPU 实际要算 50ms
end = time.time() # 测出 0.3ms ?!
0.3ms 只是"把任务塞进队列"的时间(递菜单的速度),不是计算时间(做菜的时间)。真正的 50ms 计算此刻还在队列里------这就是"哇好快"错觉的来源。
torch.cuda.synchronize() = 阻塞 CPU,直到队列里所有已提交的 GPU 任务真正完成。两个同步点缺一不可:
| 同步点 | 不加会怎样 |
|---|---|
| 开始前 | 上一个实验遗留的任务混进本次计时,起点不准 |
| 结束后 | 秒表在任务出结果之前就停了 → 荒谬的"超快" |
为什么多次运行取平均:
- 首次开销(最常见):CUDA 上下文初始化、kernel JIT 编译、显存首次分配------只发生一次;
- 缓存冷热差异:第一遍数据还在 HBM 冷着;
- 环境噪声:系统调度、其他进程、温度墙降频。
惯例:丢弃第一次,取平均(严格的 benchmark 取最小值------最小值最接近真实成本)。
更专业的工具:CUDA Event------在 GPU 时间轴上打点,直接量 GPU 自己经历了多久:
python
start_ev = torch.cuda.Event(enable_timing=True)
end_ev = torch.cuda.Event(enable_timing=True)
start_ev.record()
func()
end_ev.record()
torch.cuda.synchronize()
print(start_ev.elapsed_time(end_ev)) # 毫秒,GPU 视角的真实耗时
与估算的关系:前面用 FLOPs ÷ 有效算力得到"应该多久"(纸面值),benchmark 用同步点 + 多次平均得到"实际多久"(真值)。两者的差距就是 MFU------估算时拍的 0.5 利用率,最终靠实测来校验。
4.2 MFU (Model FLOPs Utilization)
MFU = (实际FLOP/s) / (硬件峰值FLOP/s)
- MFU >= 0.5:被认为是相当不错的性能
- MFU 接近 1.0:非常难达到,因为硬件不可能100%满负荷运转,总会有内存访问、数据传输等开销
注意:这个公式忽略了通信和系统开销,只关注纯粹的计算效率。
接下来我们在 bfloat16 上计算一下 MFU:
# 将张量转换为 bfloat16
x = x.to(torch.bfloat16)
w = w.to(torch.bfloat16)
# 测量实际性能
bf16_actual_time = time_matmul(x, w)
bf16_actual_flop_per_sec = actual_num_flops / bf16_actual_time
# 获取 bfloat16 的理论峰值
bf16_promised_flop_per_sec = get_promised_flop_per_sec(device, x.dtype)
# 计算 MFU
bf16_mfu = bf16_actual_flop_per_sec / bf16_promised_flop_per_sec
使用 bfloat16 时,actual_flop_per_sec 通常比 float32 更高,因为硬件对低精度计算进行了优化。这里的 MFU 值相当低,可能是因为硬件厂商公布的 promised_flop_per_sec 往往是过于乐观的估计(虚标)。
4.3 梯度计算
4.3.1 梯度反向传播
继续以一个简单的线性模型为例:
y = 0.5 * (x * w - 5)²
这里,x 是输入数据,w 是模型参数,y 是损失值。
前向传播代码:
x = torch.tensor([1., 2, 3]) # 输入数据,不需要计算梯度
w = torch.tensor([1., 1, 1], requires_grad=True) # 模型参数,需要计算梯度
pred_y = x @ w # 预测值(矩阵乘法)
loss = 0.5 * (pred_y - 5).pow(2) # 计算损失值
requires_grad=True 这个参数它告诉 PyTorch "请为这个张量 w 构建计算图,并在反向传播时计算它的梯度。" 对于输入数据 x,我们通常不需要它的梯度,所以不设置此标志。
loss.backward() # 触发反向传播
assert loss.grad is None # 损失值 loss 本身是一个标量,它没有"上游"的梯度,所以它的 .grad 属性是 None
assert pred_y.grad is None # 中间变量 pred_y 默认情况下也不会保存梯度,除非你显式地调用 pred_y.retain_grad()
assert x.grad is None # 因为 x 在创建时没有设置 requires_grad=True,所以它不会被追踪,其 .grad 也为 None
assert torch.equal(w.grad, torch.tensor([1, 2, 3])) # 验证模型参数 w 的梯度是否计算正确
4.3.2 计算图与自动求导
计算图是一个有向无环图(DAG),记录从输入到输出的每一步运算依赖关系:
- 节点:张量或运算操作
- 边 :数据流向
为什么需要计算图?- 自动求导 :
loss.backward()沿图反向遍历,用链式法则自动算出每个参数的梯度- 链式法则的工程实现:把复杂求导链拆成基础运算(加、乘、平方等),PyTorch 为每种运算内置求导公式,反向遍历时逐段套用
- 动态图 :每次前向传播重新构建图,支持 Python 原生控制流(
if/for),调试直观
关键细节:requires_grad=True标记需要求梯度的张量,只有它们被纳入计算图追踪- 中间变量默认不保留梯度(节省显存),只有叶子节点保留
.gradbackward()只能从标量出发,因为梯度是"损失对参数的偏导数"
反向传播代码:
调用 loss.backward() 后,PyTorch 会从 loss 开始,沿着计算图回溯,自动计算出每个需要梯度的张量的梯度。
4.3.2.1 笔记中的例子
python
x = torch.tensor([[1., 2, 3], [4, 5, 6]])
w = torch.tensor([[7., 8, 9], [10, 11, 12]], requires_grad=True)
pred_y = x @ w.t() # 矩阵乘法
loss = 0.5 * (pred_y - 5) ** 2
loss.backward()
4.3.2.2 计算图结构
x ──┐
├── @ ──→ pred_y ──→ -5 ──→ ² ──→ ×0.5 ──→ loss
w.t()┘
每一步拆解:
| 步骤 | 操作 | 输出 | 记录的信息 |
|---|---|---|---|
| ① | x @ w.t() |
pred_y |
矩阵乘法,保存输入 x 和 w |
| ② | pred_y - 5 |
中间值 | 减法,保存 pred_y |
| ③ | (...) ** 2 |
中间值 | 幂运算,保存减法结果 |
| ④ | 0.5 * (...) |
loss |
乘法,保存平方结果 |
调用 loss.backward() 时,从 loss 出发沿图反向遍历,链式法则逐段求导: |
loss = 0.5 × (pred_y - 5)²
↓ ∂loss/∂(pred_y-5)² = 0.5
(pred_y - 5)²
↓ ∂(□²)/∂□ = 2×(pred_y-5)
pred_y - 5
↓ ∂(□-5)/∂□ = 1.0
pred_y = x @ w.t()
↓ ∂(x@w.t())/∂w = x.t() @ ...(矩阵乘法求导规则)
w
链式法则串联 :每一步的局部导数相乘,最终得到 w.grad。PyTorch 为 @、-、**2、* 每种运算都内置了求导公式,反向遍历时自动套用,不需要手动推导。
w是叶子节点 :requires_grad=True,反向传播后w.grad被填充pred_y是中间变量 :pred_y.grad为None,不保留梯度以节省显存x不求梯度 :requires_grad默认为False,不纳入计算图追踪loss是标量 :才能直接调用.backward(),因为梯度是"损失对参数的偏导数"
4.3.3 计算反向传播(梯度)
if torch.cuda.is_available():
B = 16384 # Number of points
D = 32768 # Dimension
K = 8192 # Number of outputs
else:
B = 1024
D = 256
K = 64
device = get_device()
x = torch.ones(B, D, device=device)
w1 = torch.randn(D, D, device=device, requires_grad=True)
w2 = torch.randn(D, K, device=device, requires_grad=True)
Model: x --w1--> h1 --w2--> h2 -> loss
h1 = x @ w1 # 第一层:输入x乘以权重w1得到隐状态h1
h2 = h1 @ w2 # 第二层:h1乘以权重w2得到输出h2
loss = h2.pow(2).mean() # 计算损失(均方误差)
回顾一下前向传播 FLOPs 的计算,第一层 x @ w1 需要 2 * B * D * D FLOPs;第二层 h1 @ w2 需要 2 * B * D * K FLOPs,总计为 num_forward_flops = (2 * B * D * D) + (2 * B * D * K)
接下来让我们计算反向传播 FLOPs,调用 loss.backward() 后,PyTorch 需要计算以下四个梯度:
- h1.grad = d(loss) / d(h1) (中间激活值的梯度)
- h2.grad = d(loss) / d(h2) (最后一层输出的梯度)
- w1.grad = d(loss) / d(w1) (第一层权重的梯度)
- w2.grad = d(loss) / d(w2) (第二层权重的梯度)
为什么需要这四种梯度?------ 两个是"目的",两个是"工具"
| 梯度 | 角色 | 命运 |
|---|---|---|
w1.grad / w2.grad |
目的(终点站) | 训练的最终产品:优化器拿它更新权重 w -= lr * grad |
h2.grad / h1.grad |
工具(中转站) | 没人直接用它,存在的唯一理由是:没有它们就算不出权重梯度 |
为什么绕不开激活梯度?------ 链式法则的传递路径
loss 只直接依赖 h 2 h_2 h2,它压根不知道 h 1 h_1 h1 和 w 1 w_1 w1 的存在。想知道" w 1 w_1 w1 动一动 loss 会怎么变",只能沿前向的计算路径倒着推 :
Model: x --w1--> h1 --w2--> h2 -> loss
对 w 2 w_2 w2(链短,只穿过一层):
∂ l o s s ∂ w 2 = ∂ l o s s ∂ h 2 ⋅ ∂ h 2 ∂ w 2 \frac{\partial loss}{\partial w_2} = \frac{\partial loss}{\partial h_2} \cdot \frac{\partial h_2}{\partial w_2} ∂w2∂loss=∂h2∂loss⋅∂w2∂h2
对 w 1 w_1 w1(链长,要穿过中间激活):
∂ l o s s ∂ w 1 = ∂ l o s s ∂ h 2 ⋅ ∂ h 2 ∂ h 1 ⋅ ∂ h 1 ∂ w 1 \frac{\partial loss}{\partial w_1} = \frac{\partial loss}{\partial h_2} \cdot \frac{\partial h_2}{\partial h_1} \cdot \frac{\partial h_1}{\partial w_1} ∂w1∂loss=∂h2∂loss⋅∂h1∂h2⋅∂w1∂h1
规律:权重离 loss 越深、链越长、需要的激活梯度越多;
且深层的链恰好以浅层的激活梯度为前缀------这就是"梯度逐层回传、沿途复用"的反向传播。
矩阵乘法的求导规则给出每层的公式:
- h1.grad = d(loss) / d(h1) (中间激活值的梯度)
- h2.grad = d(loss) / d(h2) (最后一层输出的梯度)
- w1.grad = d(loss) / d(w1) (第一层权重的梯度)
- w2.grad = d(loss) / d(w2) (第二层权重的梯度)
w 2 . g r a d = h 1 T ⋅ h 2 . g r a d w 1 . g r a d = x T ⋅ h 1 . g r a d w_2.grad = h_1^T \cdot h_2.grad \qquad w_1.grad = x^T \cdot h_1.grad w2.grad=h1T⋅h2.gradw1.grad=xT⋅h1.grad
可见 h 2 . g r a d h_2.grad h2.grad 是算 w 2 . g r a d w_2.grad w2.grad 的原料, h 1 . g r a d h_1.grad h1.grad( = h 2 . g r a d @ w 2 T = h_2.grad @ w_2^T =h2.grad@w2T)是算 w 1 . g r a d w_1.grad w1.grad 的原料------激活梯度就是链式法则的接力棒。
反向传播每经过一层,流入的激活梯度干两件事:
- 就地消费:和本层输入相乘,产出本层的权重梯度;
- 继续上传 :穿过本层权重(乘 w T w^T wT),变成更上一层的激活梯度。
完整反向流程:
loss
└─▶ h2.grad = ∂loss/∂h2 ← 起点:损失对输出的敏感度
├─▶ w2.grad = h1ᵀ @ h2.grad ← 【产品】第二层权重的更新依据
└─▶ h1.grad = h2.grad @ w2ᵀ ← 中转:误差传回第一层
└─▶ w1.grad = xᵀ @ h1.grad ← 【产品】第一层权重的更新依据
呼应两个前文主题 :
① 为什么训练要存激活值(吃显存) ------看 w 2 . g r a d = h 1 T @ h 2 . g r a d w_2.grad = h_1^T @ h_2.grad w2.grad=h1T@h2.grad:需要前向时的 h 1 h_1 h1!所以中间结果必须留在显存里等反向取用;推理不用存,显存小得多。
② 为什么总 FLOPs 是 6ND(前向 2 + 反向 4) ------注意流程里每层做了两次矩阵乘法:一次产权重梯度、一次传激活梯度,所以反向 ≈ 前向的两倍,合计 2+4=6。
4.4 优化器 (Optimizer)
4.4.1 常用优化器介绍
- SGD (随机梯度下降):最基础的优化器,直接用学习率乘以梯度更新参数。
- Momentum (动量法):在 SGD 基础上增加了一个"动量"项,即梯度的指数移动平均值,有助于加速收敛并减少震荡。
- AdaGrad:根据历史梯度的平方值来调整每个参数的学习率,对稀疏特征更友好。
- RMSProp:对 AdaGrad 的改进,使用梯度平方的指数加权平均代替简单累加,避免学习率过早衰减。
- Adam:融合了 RMSProp 和 Momentum 的思想,是目前最流行的优化器。
我们以 AdaGrad 为例(虽然现在常用 AdamW,但原理类似)。优化器不仅要更新参数,还要记住每个参数的历史梯度信息(状态)。
class AdaGrad(torch.optim.Optimizer):
def step(self):
for group in self.param_groups:
for p in group['params']:
grad = p.grad.data
# 获取状态(梯度平方和)
state = self.state[p]
if 'sum_squared_grad' not in state:
state['sum_squared_grad'] = torch.zeros_like(p.data)
# 更新状态:累加梯度的平方
state['sum_squared_grad'] += grad ** 2
# 更新参数:除以 根号下(状态)
std = state['sum_squared_grad'].sqrt() + 1e-10
p.data -= group['lr'] * grad / std
4.4.2 优化器的使用
在 PyTorch 中实例化和使用一个 AdaGrad 优化器:
# 实例化优化器
optimizer = AdaGrad(model.parameters(), lr=0.01) # model.parameters():将模型中所有可学习的参数传递给优化器
# 计算梯度
loss.backward() # 计算损失函数对所有参数的梯度
# 执行一步更新
optimizer.step() # 根据梯度和优化器内部状态,更新模型参数
4.5 资源核算
4.5.1 内存占用分析
对于一个深度线性模型,总内存需求由四部分组成:
-
参数 (Parameters):模型中所有可学习权重的数量。
-
激活值 (Activations):前向传播过程中产生的中间结果,需要保存下来用于反向传播。
-
梯度 (Gradients):反向传播计算出的梯度,其数量与参数相同。
-
优化器状态 (Optimizer States):优化器维护的额外状态信息(如 AdaGrad 的 g2),其数量也与参数相同。
假设所有数据都使用 float32 格式(每个元素占 4 字节),则总内存为:total_memory = 4 * (num_parameters + num_activations + num_gradients + num_optimizer_states)
4.5.2 计算量(FLOPs)分析
总 FLOPs 约等于 6 × (数据量) × (参数量)。因此,对于一个训练步骤,其计算量为:
flops = 6 * B * num_parameters
4.6 混合精度训练 (Mixed Precision)
4.6.1 混合精度训练概述
问题 :fp32 能稳定训练但内存太大;fp16/bf16 省内存但有数值不稳定风险。如何在"高精度的稳定性"和"低精度的效率"之间取得平衡?
解决方案 :混合精度训练 (Mixed Precision Training, 2017)
解决方案是采用混合精度策略。默认使用 float32,确保关键部分的计算精度。在可能的情况下使用 {bfloat16, fp8},利用其高效的内存和计算特性。如下给出一个经典的混合精度训练方案:
- 前向传播(Forward Pass):使用 bfloat16 或 fp8。这包括所有中间激活值(activations)。因为激活值通常不需要极高的精度,使用低精度可以显著节省内存。
- 其余部分:使用 float32。这包括模型参数(parameters)、梯度(gradients)以及优化器状态(optimizer states)。这些是训练的核心,需要更高的精度来保证数值稳定性和收敛性。
核心思想:将低精度用于"消耗大但对精度要求不高"的部分(激活值),将高精度用于"对精度敏感"的部分(参数和梯度)。
4.6.2 自动实现混合精度训练的工具
这里我们介绍两个主要的工具库,它们可以自动化地实现混合精度训练:
- PyTorch 的 AMP 库 (Automatic Mixed Precision):PyTorch 提供了自动混合精度(AMP)库,会自动将安全的操作(如矩阵乘法)转为 bf16,将危险操作(如 exp、softmax)保留为 fp32:
python
with torch.amp.autocast("cuda", dtype=torch.bfloat16):
x = torch.zeros(4, 8) # 自动以 bf16 创建
- NVIDIA Transformer Engine:这是一个专门针对 Transformer 模型优化的库,它支持在矩阵乘法等核心操作中使用 FP8 精度。目标是实现全链路的 FP8 训练,即在整个训练过程中都使用 FP8,以达到极致的性能和效率
4.7 算术强度与 Roofline 分析
简化 GPU 的工作模型:
- 从 HBM(高带宽内存)发送输入到计算核心
- 执行计算
- **将输出从计算核心送回 HBM
4.8 ReLU 的算术强度分析
以一个简单的 ReLU 操作为例(1024 × 1024 维的 bf16 向量):
内存移动(Bytes):
- 读入 x:2 × n(bf16 是 2 字节/元素)
- 写出 y:2 × n
- 总计:4n
计算量(FLOPs): - n 次比较(max(x, 0))
- 共 n FLOPs
通信时间 = 4n / (3.35 × 10¹²) ≈ 1.2 × 10⁻⁶ 秒 计算时间 = n / (989.5 × 10¹²) ≈ 1.0 × 10⁻⁹ 秒
这里有一个重要假设:通信和计算可以完美重叠(overlap) 。在理想情况下,数据到达后立即开始计算,计算的同时下一批数据已在传输。因此总时间 = max(通信时间, 计算时间),而不是两者之和。
在这个例子中,通信时间远超计算时间 ------ReLU 是典型的 memory-bound(内存受限)操作。
4.9 算术强度的定义
为了避免每次都要计算两个时间再比较,引入算术强度(Arithmetic Intensity):
加速器强度 = FLOP/s / Bytes/s → H100 约 295 FLOP/byte
算术强度 = FLOPs / Bytes → 该操作每搬运 1 字节能做多少 FLOP
- 算术强度 < 加速器强度 → Memory-bound(瓶颈在数据传输)
- 算术强度 > 加速器强度 → Compute-bound(瓶颈在计算)
"对于 H100,加速器强度大约是 295。这个数字值得记一下------对于 bf16 来说,每搬运 1 字节你需要做大约 300 次浮点运算才能摆脱内存瓶颈。"
4.10 Roofline 图

Roofline 图直观地展示了算术强度与性能的关系:
- X 轴:算术强度(每个"切片"对应一个特定算法)
- Y 轴:实际达到的 FLOP/s
- 每条分段线性曲线:一个特定的硬件平台(H100、B200 等)
- 转折点(kink):该硬件的加速器强度------转折点左侧是 memory-bound 区域(斜率上升),右侧是 compute-bound 区域(水平天花板)