CS336课程笔记:lecture2 pytorch手把手搭建 LLM
CS336课程笔记:lecture2 pytorch手把手搭建 LLM
目录
- 0. 本讲概览与学习目标
- 1. PyTorch 基础与资源效率动机
- 2. 内存资源与张量 (tensor) 基础
- 3. 计算量与 FLOPs 估算
- 4. 梯度、模型与参数统计
- 5. 训练循环与工程实践
- 6. 数值精度、混合精度训练与硬件对比
0. 本讲概览与学习目标
本讲是本课程的第二讲,延续上一节关于 语言模型 (Language Model) 与 从零实现 的总览,开始真正动手用 PyTorch 搭建模型,并围绕资源效率(时间与显存)做系统的分析与估算。
核心问题:
- 在给定 GPU 资源(如 A100/H100 数量、显存大小、FLOPs 能力)的情况下:
- 能训练多大的模型?(参数规模)
- 要花多长时间?(训练时长)
- 如何在 PyTorch 中写出既正确又高效的代码?
解决方案路径:
-
步骤 A:内存与张量基础
理解张量数据类型、维度与在 GPU 上的存储方式,学会按字节数估算显存占用,并从中得出“参数最多能有多少”、“激活会占多少”等结论。 -
步骤 B:计算量与 FLOPs 分析
分析矩阵乘法、线性层等核心算子的 FLOPs,推导类似 6 N P 6NP 6NP 这种用于估算 LLM 训练开销的经验公式,并理解 MFU (Model FLOPs Utilization) 的含义。 -
步骤 C:从模型到训练循环与数值精度
在 PyTorch 中从张量出发搭建模块与模型,写出完整训练循环,理解优化器、梯度回传与随机性控制;进一步讨论 float32、bfloat16、fp8 等不同精度下的性能差异和 混合精度训练 (Mixed Precision Training) 实践。
学习目标:
- 能够根据 GPU 数量、显存和带宽 粗略推算能训练的最大模型规模和训练时间。
- 熟练掌握 PyTorch 张量、模块、优化器和训练循环 的基本用法。
- 理解 内存占用 = 参数 + 梯度 + 优化器状态 + 激活 的分解方式。
- 能用 FLOPs 的视角理解 矩阵乘法、线性层、梯度计算 的计算量。
- 了解 float32 / bfloat16 / fp8 的精度与性能权衡,知道为什么要使用 自动混合精度 (AMP) 与专用库提高效率。
1. PyTorch 基础与资源效率动机
1.1 课程承接与本讲定位
讲师首先回顾上一讲内容:
- 上一讲主要介绍了:
- 语言模型总体框架;
- 从零实现 (build from scratch) 的动机;
- 分词 (tokenization) 和作业中 tokenizer 的实现(作业第一部分)。
- 本讲将从 动手实现 的角度出发:
- 使用 PyTorch 搭建模型;
- 理解训练一个大模型需要的 计算资源与内存资源;
- 强调:效率 (efficiency) 与 算力成本 和 金钱成本 直接相关。
本讲不会详细推导 Transformer 架构 本身,而是:
- 只用更简单的模型做示范;
- 把重点放在:
- PyTorch 原语 (primitives) 的使用;
- 资源核算 (resource accounting) 的思维方式。
学生需要记住:
- 机械层面 (mechanics):如何用 PyTorch 表示张量、搭建模块、写训练循环;
- 思维层面 (mindset):每写一段代码,都要有大致的 显存与 FLOPs 概念,知道这段代码“贵不贵”。
1.2 通过“纸上算一算”激发资源意识
讲师一开始给出两个“纸上算一算”的问题,用来说明为什么资源核算很重要。
问题 1:训练 70B 参数 Transformer 需要多久?
问题设定:
- 模型:70B 参数的致密 (dense) Transformer;
- 数据:15T tokens;
- 硬件:1024 张 A100 GPU;
- 目标:估算 训练完一次的时间。
解决思路:
-
先估算 总 FLOPs:
- 经验公式:
Total FLOPs ≈ 6 × N params × N tokens \text{Total FLOPs} \approx 6 \times N_\text{params} \times N_\text{tokens} Total FLOPs≈6×Nparams×Ntokens - 其中:
- N params N_\text{params} Nparams:模型参数总数,这里是 70 B 70\,\text{B} 70B;
- N tokens N_\text{tokens} Ntokens:训练用的 token 总数,这里是 15 T 15\,\text{T} 15T;
- 系数 6 来自一次前向和反向在 Transformer 中的典型 FLOPs 统计,本讲后面会解释其来源。
- 经验公式:
-
再根据 GPU 厂商给出的 单卡理论峰值 FLOPs 与设置的 MFU,计算集群每天能给出的 FLOPs 总量:
- 查表得到某类型 A100 的理论 FLOPs(与数据类型相关);
- 选择一个合理的 MFU (Model FLOPs Utilization),例如设置为 0.5 0.5 0.5;
- 计算:
FLOPs/day = GPU 理论 FLOPs × GPU 数量 × MFU × 86400 秒 \text{FLOPs/day} = \text{GPU 理论 FLOPs} \times \text{GPU 数量} \times \text{MFU} \times 86400 \;\text{秒} FLOPs/day=GPU 理论 FLOPs×GPU 数量×MFU×86400秒
-
用“需要的 FLOPs 总量”除以“每天能给的 FLOPs”:
训练天数 = Total FLOPs FLOPs/day \text{训练天数} = \frac{\text{Total FLOPs}}{\text{FLOPs/day}} 训练天数=FLOPs/dayTotal FLOPs
结果:
- 讲师给出的近似结果是 约 144 天。
需要记住的要点:
- 只要知道 参数量、token 数、硬件算力与目标 MFU,就可以很快得到一个数量级正确的训练时间估计。
- 系数 6 6 6 的来源与 Transformer 的结构有关,但在工程上,“先用经验公式” 就能帮助做预算和规划。
问题 2:单卡 A100 上能放多大的模型?
问题设定:
- 硬件:单张 A100 GPU,显存 80GB;
- 优化器:AdamW;
- 假设暂时不考虑激活占用,只关心 参数 + 梯度 + 优化器状态;
- 目标:估算 最多可训练的参数量。
解决思路:
-
记住一个经验数:
- 对 AdamW:
- 参数本身:1 份;
- 梯度:1 份;
- 一阶矩动量 m m m:1 份;
- 二阶矩动量 v v v:1 份;
- 如果这些都用 float32 (4 bytes) 存储,则每个参数大约需要 16 字节。
- 对 AdamW:
-
用显存总量除以每个参数所需的字节数:
N params ≈ 80 GB 16 bytes N_\text{params} \approx \frac{80\,\text{GB}}{16\,\text{bytes}} Nparams≈16bytes80GB -
粗略结果:
- 得到的数量级约为 40B 参数。
讲师特别强调:
- 这个估算 还没有算上激活 (activations),而激活占用和 batch size、序列长度 强相关,在作业中会很关键;
- 即便如此,这个粗略计算已经足以让我们在设计模型时有大致概念,不至于“随手一写就 OOM”。
学生需要记住:
- 内存核算 是训练大模型的第一步;
- 以 AdamW 为例,“16 bytes / 参数” 是一个非常重要的经验数,后面会在混合精度中看到如何降低这个数字。
1.3 本讲结构总览
讲师把本讲总结为三个主线:
-
Memory accounting(内存核算):
- 从 张量 (tensor) 基础 开始,理解数据类型和在 GPU 上的存储;
- 学会估算参数、优化器状态、激活的显存消耗;
- 为后续的 模型规模上限 与 batch size 选择 提供依据。
-
Compute accounting(计算核算):
- 从简单的矩阵乘法开始,统计 FLOPs;
- 推到线性层、残差连接、注意力等操作的大致 FLOPs;
- 连接到一开始的 6 N P 6NP 6NP 公式与 MFU 概念。
-
PyTorch primitives + training loop(PyTorch 原语与训练循环):
- 实际在 PyTorch 中写代码:
- 张量在 GPU 上的创建与操作;
- 使用
torch.nn.Module组织模型; - 构造优化器与训练循环;
- 结合 随机性、checkpoint、混合精度 等工程实践,构建可复现实验。
- 实际在 PyTorch 中写代码:
这一段的核心信息:
- 本讲不是在教“怎么调参”,而是在教“如何把资源算清楚”;
- 这种从底层算起的习惯,会贯穿后续所有关于大模型训练的内容。
2. 内存资源与张量 (tensor) 基础
本部分主要对应 PPT 中的 “Memory accounting / tensors basics / tensors_memory / tensors_on_gpus / tensor_operations / tensor_einops” 等内容。
2.1 张量是深度学习的基本载体
讲师明确指出:
- 张量 (tensor) 是深度学习中所有数据的统一表示形式:
- 训练数据(样本、batch);
- 模型参数 (parameters);
- 中间激活 (activations);
- 梯度 (gradients) 和优化器状态 (optimizer states)。
直觉理解:
- 可以把张量看成是 多维数组:标量、向量、矩阵都只是不同维度的特殊情况;
- PyTorch 用张量对象封装了:
- 数据本身(存放在 CPU/GPU 内存或显存);
- 数据类型 dtype(如
float32/bfloat16/int64等); - 设备信息 device(如
cpu/cuda:0)。
学生需要记住:
- 任何关于“内存占用”的问题,归根结底都是在问:
占用字节数 = 元素个数 × 每个元素的字节数 \text{占用字节数} = \text{元素个数} \times \text{每个元素的字节数} 占用字节数=元素个数×每个元素的字节数
2.2 张量的内存占用:tensors_memory()
讲师用简单示例说明如何在 PyTorch 中估算张量占用:
- 若有一个形状为
(
m
,
n
)
(m, n)
(m,n) 的矩阵
X
X
X,数据类型为
float32:- 元素个数: m × n m \times n m×n;
- 每个元素 4 字节;
- 总占用:
bytes ( X ) = m × n × 4 \text{bytes}(X) = m \times n \times 4 bytes(X)=m×n×4
更一般地,若张量形状为 ( d 1 , d 2 , … , d k ) (d_1, d_2, \dots, d_k) (d1,d2,…,dk),则:
bytes ( X ) = ( ∏ i = 1 k d i ) × bytes_per_element \text{bytes}(X) = \left(\prod_{i=1}^k d_i\right) \times \text{bytes\_per\_element} bytes(X)=(i=1∏kdi)×bytes_per_element
其中:
- d i d_i di:第 i i i 维的长度;
bytes_per_element:由 dtype 决定,例如:float32:4 bytes;bfloat16/float16:2 bytes;int8:1 byte。
课堂中强调:
- 参数张量、梯度张量 与 优化器状态张量 都可以用同样的方式估算显存;
- 内存核算时,要明确区分:
- 静态占用(参数、优化器状态)——与 batch size 无关;
- 动态占用(激活)——与 batch size、序列长度、模型结构强相关。
学生需要记住的结论:
- 在做“最大可训练模型规模”估算时,通常先算 参数 + 优化器状态:
- 这是显存的硬下限;
- 再预留出一部分空间给激活和临时张量。
2.3 张量与 GPU:tensors_on_gpus()
为了真正利用 GPU 的大规模并行计算能力,张量需要被移动到 GPU 上:
- 典型代码:
x = x.to("cuda")或x = x.cuda();- 模型同样需要
.to(device)。
关键点:
- 张量所在的设备必须一致,算子才能正常工作:
- 若输入张量在 CPU、权重张量在 GPU,则会报错或触发隐式数据拷贝;
- 隐式数据拷贝既慢又难以察觉,应该主动管理。
讲师强调的实践建议:
- 在代码一开始就确定一个
device变量:- 例如
device = torch.device("cuda" if torch.cuda.is_available() else "cpu"); - 所有张量和模型都统一使用这个 device。
- 例如
学生需要记住:
- 在做 FLOPs 估算与时间测量 时,必须确保张量真正处于 GPU 上,避免“在 CPU 上算”的假象;
- 在大模型场景下,任何不必要的 CPU-GPU 拷贝都要尽量消除。
2.4 张量操作与内存模式:tensor_operations()
本节讲师展示了几类典型张量操作:
- 逐元素 (elementwise) 运算:如加法、减法、
relu等; - 矩阵乘法 (matmul) 与 线性层;
- 广播 (broadcasting);
- 视图变换 (view/reshape/transpose)。
就内存角度来看:
- 视图变换 通常不复制数据,只改变张量的 shape 和 stride,几乎不增加额外内存;
- 逐元素运算 通常会产生新的张量,若不小心保存中间结果可能额外占用显存;
- 矩阵乘法 和 卷积 既是主要的 FLOPs 大户,也是重要的激活来源。
学生需要记住:
- 在写模型时尽量利用 in-place 操作 或合理的重用,减少不必要的中间张量;
- 但同时要注意:某些 in-place 操作会与自动求导系统冲突,需要谨慎使用(本讲在梯度部分会再提)。
2.5 使用 einops 进行张量重排:tensor_einops()
PPT 中提到 einops 库和 einops.rearrange / einsum / reduce 等函数,用来做:
- 形状变换、维度重排;
- 便捷地表达某些矩阵运算或批处理操作。
在内存与效率角度:
einsum可以以较清晰的语法表示复杂的张量乘法;- 很多时候,合理的
einsum写法能让我们更容易地数出 FLOPs。
学生需要记住:
- 对于大模型中的多维张量(例如
[batch, seq, head, dim]),einops是非常实用的工具; - 但无论使用什么接口,本质都是某种矩阵/张量乘法,在计算量和内存占用上不会“凭空变少”。
3. 计算量与 FLOPs 估算
这一部分对应 PPT 中的 compute accounting / tensor_operations_flops / gradients_flops 等内容,是本讲的核心之一。
3.1 从矩阵乘法开始的 FLOPs 直觉
讲师首先讨论了 矩阵乘法 (matrix multiplication) 的 FLOPs:
- 若有矩阵
A
∈
R
m
×
k
A \in \mathbb{R}^{m \times k}
A∈Rm×k、
B
∈
R
k
×
n
B \in \mathbb{R}^{k \times n}
B∈Rk×n,则:
- 结果矩阵 C = A B C = A B C=AB 的形状为 ( m , n ) (m, n) (m,n);
- 计算每个元素
C
i
j
C_{ij}
Cij 需要:
- k k k 次乘法;
- k − 1 k-1 k−1 次加法;
- 整体 FLOPs 量级为:
FLOPs ( C = A B ) = O ( m k n ) \text{FLOPs}(C = A B) = O(m k n) FLOPs(C=AB)=O(mkn)
在实际估算中,常用简化:
- 忽略常数,只记作 2 m k n 2mkn 2mkn 或 m k n mkn mkn 这样的量级;
- 更重要的是:
- 认清楚 矩阵乘法 FLOPs 与三个维度线性相关,即只要你把一个维度乘以 2,计算量大致也乘以 2。
学生需要记住:
- 在深度学习中,大部分 FLOPs 都来自矩阵乘法或卷积;
- 只要你知道输入输出维度,就可以快速估计它们的 FLOPs;
- 逐元素运算(如
relu、加减等)相对便宜,在大模型中常常不是瓶颈。
3.2 线性层与多层结构的 FLOPs
考虑一个线性层:
- 输入:形状为 ( B , D in ) (B, D_\text{in}) (B,Din) 的张量;
- 权重:形状为 ( D in , D out ) (D_\text{in}, D_\text{out}) (Din,Dout);
- 输出:形状为 ( B , D out ) (B, D_\text{out}) (B,Dout);
- 本质是一次矩阵乘法 X W XW XW,FLOPs 约为:
FLOPs linear ≈ 2 B D in D out \text{FLOPs}_\text{linear} \approx 2 B D_\text{in} D_\text{out} FLOPslinear≈2BDinDout
其中:
- B B B:batch size;
- D in D_\text{in} Din:输入维度;
- D out D_\text{out} Dout:输出维度。
当把多个线性层堆叠成深度网络时:
- 总 FLOPs 大约是 各层 FLOPs 的求和;
- 若各层维度相近,常用简化:
- 假设每层参数量约为 P P P,深度为 L L L,则总参数量为 L P LP LP;
- 每次前向传播 FLOPs 与参数量近似同阶。
学生要记住:
- 参数量与 FLOPs 一般是同一个量级(至少在全连接或 Transformer 中);
- 这就是为什么可以用 “ 6 N params N tokens 6 N_\text{params} N_\text{tokens} 6NparamsNtokens” 的形式来估算训练总 FLOPs。
3.3 训练过程中的 FLOPs:前向 + 反向
在训练中,每一步需要:
- 前向传播 (forward pass):计算模型输出与损失;
- 反向传播 (backward pass):计算梯度;
- 参数更新 (optimizer step):根据梯度更新参数。
FLOPs 角度:
-
对大多数网络结构:
- 反向传播 FLOPs 约为前向的 2 倍;
- 参数更新 FLOPs 相对较小,可以忽略或作为常数因子;
-
因此一轮训练(一次参数更新)总 FLOPs 大约是:
FLOPs train per token ≈ 3 × FLOPs forward per token \text{FLOPs}_\text{train per token} \approx 3 \times \text{FLOPs}_\text{forward per token} FLOPstrain per token≈3×FLOPsforward per token
结合 Transformer 的结构和多头注意力等操作,最终可以得到类似:
Total FLOPs ≈ 6 × N params × N tokens \text{Total FLOPs} \approx 6 \times N_\text{params} \times N_\text{tokens} Total FLOPs≈6×Nparams×Ntokens
讲师在这里强调:
- 精确常数因子并不重要,重要的是 会算数量级;
- 在实际工程中,我们常以这种经验公式为起点,然后再通过 profile 工具 进一步精细分析。
学生需要记住:
- 训练开销与参数量和 token 数成正比;
- 把 token 数翻倍,训练时间几乎也会翻倍;
- 把模型参数翻倍,如果其他条件不变,训练时间也会近似翻倍。
3.4 用 PyTorch 实测 FLOPs 与时间:tensor_operations_flops()
在 PPT 中,讲师展示了一个实验流程,用 PyTorch 时间函数来验证计算:
-
定义一个矩阵乘法函数
time_matmul(x, w):- 输入张量
x和权重矩阵w放在 GPU 上; - 使用
torch.cuda.synchronize()确保计时时间准确; - 返回一次 matmul 的实际耗时
actual_time。
- 输入张量
-
预先用公式计算该 matmul 的理论 FLOPs:
actual_num_flops = 2 m k n \text{actual\_num\_flops} = 2 m k n actual_num_flops=2mkn
-
用公式:
actual_flop_per_sec = actual_num_flops actual_time \text{actual\_flop\_per\_sec} = \frac{\text{actual\_num\_flops}}{\text{actual\_time}} actual_flop_per_sec=actual_timeactual_num_flops
-
对比 GPU 厂商给出的 理论峰值 FLOPs:
- 通过类似
get_promised_flop_per_sec(device, x.dtype)的辅助函数查表; - 得到
promised_flop_per_sec。
- 通过类似
-
计算 MFU (Model FLOPs Utilization):
MFU = actual_flop_per_sec promised_flop_per_sec \text{MFU} = \frac{\text{actual\_flop\_per\_sec}}{\text{promised\_flop\_per\_sec}} MFU=promised_flop_per_secactual_flop_per_sec
讲师指出:
- 在一个设计合理的矩阵乘法 benchmark 中,当矩阵足够大,使得 GPU 利用充分时:
- MFU 大于 0.5 通常就被认为相当不错;
- 如果 MFU 很低,就说明程序中有大量开销没有用在真正的 matmul 上(比如数据搬运、 kernel 启动过多、维度不合适等)。
学生需要记住:
- MFU = 实际 FLOPs / 理论 FLOPs,越高越好;
- 在做大规模训练时,追踪 MFU 可以帮助诊断性能问题;
- 理想情况下,我们希望大部分时间都花在大矩阵乘法上,而不是在数据搬运或小算子上浪费。
4. 梯度、模型与参数统计
本部分对应 PPT 中的 gradients_basics / gradients_flops / module_parameters / custom_model 等内容。
4.1 梯度与自动求导:gradients_basics()
讲师简要回顾了 PyTorch 的 自动求导 (autograd) 机制:
- 每个张量可以设置
requires_grad=True; - PyTorch 在前向传播时构建计算图;
- 在调用
loss.backward()时,自动沿着计算图反向传播,计算出所有叶子张量的梯度。
关键直觉:
- 在内存层面,自动求导需要 保存前向传播中的中间激活,以便在反向时使用;
- 这就是为什么 激活显存 随着网络深度和 batch size 增长,常常比参数显存更大。
学生需要记住:
- 对于每一层:
- 要计算梯度,就必须在反向时用到前向的输入或中间量;
- 这意味着:训练时的显存占用 > 推理时的显存占用;
- 许多“节省显存”的技巧(如 checkpointing)本质上是在 用时间换空间:
- 不保存所有激活,而是在反向时重新计算一部分前向。
4.2 梯度的 FLOPs:gradients_flops()
在计算 FLOPs 时,不能只算前向:
- 对于一个线性层
y
=
W
x
y = Wx
y=Wx:
- 前向:一次矩阵乘法;
- 反向:需要计算:
- w.r.t. W W W 的梯度:与 x x x 和 ∂ L / ∂ y \partial L/\partial y ∂L/∂y 相关的矩阵乘法;
- w.r.t. x x x 的梯度:另一次矩阵乘法。
结果是:
- 反向 FLOPs 大约是前向的两倍;
- 因此,整体训练 FLOPs 是前向的约 3 倍(前向 1,反向 2)。
学生需要记住:
- 这也是 “训练比推理贵很多” 的根本原因;
- 在估算大模型训练成本时,不能只看一次前向,而要考虑完整的“前向 + 反向 + 更新”。
4.3 模块与模型:module_parameters() 与 custom_model()
讲师接着展示如何在 PyTorch 中构建模型:
- 使用
torch.nn.Module作为基类; - 在
__init__中定义子模块和参数; - 在
forward方法中使用这些子模块和张量操作。
重要接口:
-
model.parameters():返回模型中的所有参数张量迭代器; -
可以通过遍历它们求总参数量:
N params = ∑ p ∈ params ∏ i d i ( p ) N_\text{params} = \sum_{p \in \text{params}} \prod_i d_i(p) Nparams=p∈params∑i∏di(p)
其中 d i ( p ) d_i(p) di(p) 是参数 p p p 的每个维度长度。
讲师强调:
- 在大模型场景下,随时知道模型的总参数量 是基本要求;
- 这不仅决定了训练成本,也与 推理延迟、显存占用 等密切相关。
学生需要记住:
- 编写模型类时,要习惯性地写一个小工具函数,打印参数规模和大致显存占用;
- 只要有了
N_\text{params},就可以直接代入前面提到的 6 N params N tokens 6 N_\text{params} N_\text{tokens} 6NparamsNtokens 公式做成本估算。
5. 训练循环与工程实践
本部分对应 PPT 中的 train_loop / note_about_randomness / data_loading / optimizer / checkpointing / mixed precision training 等主题。
5.1 标准训练循环结构:train_loop()
讲师给出了一个典型的训练循环骨架:
- 迭代数据 loader:
- 从数据集中取出一个 batch。
- 将数据移动到
device:x = x.to(device),y = y.to(device)。
- 前向传播:
logits = model(x);- 计算损失
loss = criterion(logits, y)。
- 反向传播:
optimizer.zero_grad();loss.backward()。
- 参数更新:
optimizer.step()。
- 记录指标与日志:
- loss、学习率、梯度范数、MFU 等。
讲师特别强调几点实践细节:
- 零梯度方式:
- 推荐使用
optimizer.zero_grad(set_to_none=True),以减少内存访问;
- 推荐使用
- 梯度裁剪:
- 对某些模型(尤其是 RNN、Transformer),需要按范数裁剪梯度,防止梯度爆炸;
- 日志与监控:
- 对大规模训练,必须持续监控 loss 曲线和硬件利用率,以早发现问题。
学生需要记住:
- 虽然训练循环看似模板化,但其中每一步都与 性能和稳定性 强相关;
- 能够写出一个干净、易读且易于 profile 的训练循环,是从“会用 PyTorch”到“会做大模型工程”的重要一步。
5.2 随机性与可复现:note_about_randomness()
在大模型训练中,可复现性非常重要。讲师提到:
- 源头包括:
- 数据打乱 (dataloader 的 shuffle);
- 权重初始化;
- dropout 等随机层;
- GPU 上的某些非确定性算子。
常见做法:
- 设置随机种子:
torch.manual_seed(seed),random.seed(seed),numpy.random.seed(seed);
- 配置 deterministic 选项:
- 对于某些算子,可以设置
torch.use_deterministic_algorithms(True),但会牺牲性能;
- 对于某些算子,可以设置
- 保存和恢复训练状态:
- 包括模型权重、优化器状态、学习率调度器状态等。
学生需要记住:
- 对于大规模实验,记录随机种子与所有重要超参数 非常关键;
- 即便做不到完全 bit-level 的复现,至少要能在宏观行为上接近。
5.3 数据加载:data_loading()
讲师简要提到数据加载的几个要点:
- 使用
DataLoader与Dataset抽象; - 合理设置
num_workers和pin_memory,避免数据加载成为瓶颈; - 对于超大规模语料:
- 常使用预处理后的二进制格式和 mmap 方式;
- 结合流水线与缓存机制,确保 GPU 不会因等待数据而闲置。
学生需要记住:
- 在大模型训练中,I/O 瓶颈会严重浪费 GPU 资源;
- 即使模型本身写得很高效,如果数据加载跟不上,整体 MFU 仍然会很低。
5.4 优化器与状态:optimizer()
讲师重点以 AdamW 为例说明优化器对显存的影响:
-
对每个参数,需要额外保存:
- 梯度;
- 一阶矩动量 m m m;
- 二阶矩动量 v v v;
-
若全部使用
float32,且参数本身也是float32,则:bytes per param ≈ 4 ( param ) + 4 ( grad ) + 4 ( m ) + 4 ( v ) = 16 bytes \text{bytes per param} \approx 4(\text{param}) + 4(\text{grad}) + 4(m) + 4(v) = 16\;\text{bytes} bytes per param≈4(param)+4(grad)+4(m)+4(v)=16bytes
这正是前面“40B 参数”推算的依据。
讲师也提及:
- 若使用更高级的优化器或额外的状态(如动量、二阶信息),显存占用会进一步增大;
- 可以通过:
- 更轻量的优化器;
- 分布式优化器(如 ZeRO)
来缓解,但这些超出本讲范围。
学生需要记住:
- 优化器不仅影响收敛速度,还直接决定显存占用;
- 在估算显存时,必须把优化器状态考虑进去,而不仅仅是参数本身。
5.5 Checkpoint 与训练中断恢复:checkpointing()
PPT 提到 “checkpoint in” 的字样,说明本讲也对 checkpoint 做了提醒:
- 定期保存:
- 模型权重;
- 优化器状态;
- 当前学习率、步数、随机种子状态等。
- 目的:
- 训练过程中若出现故障,可以从最近 checkpoint 继续;
- 便于在不同阶段分析模型性能;
- 在多实验对比中复用已有 checkpoint。
学生需要记住:
- 对于动辄训练数周的大模型,不做 checkpoint 是不可接受的;
- Checkpoint 文件本身也很大,需要合理规划存储与清理策略。
6. 数值精度、混合精度训练与硬件对比
本部分对应 PPT 中后半段关于 不同浮点精度 (float32 / bfloat16 / fp8) 与 混合精度训练 (mixed precision training) 的讨论,以及对 A100 与 H100 性能对比的表格与图示。
6.1 不同数值精度与性能权衡
讲师先回顾了几种常见数值格式:
-
float32:
- 32 位浮点数,约 7 位十进制有效数字;
- 动态范围和精度都较好,是传统深度学习默认类型;
- 单元素 4 字节,显存与 FLOPs 开销最大。
-
bfloat16:
- 16 位浮点数,但保留了与 float32 相同位数的指数部分;
- 动态范围接近 float32,但有效数字较少;
- 单元素 2 字节,显存占用减半,很多 GPU 对其有专门加速。
-
fp8:
- 8 位浮点格式,有多种具体方案;
- 显存极小,硬件实现复杂,需要更精细的缩放与校准;
- 主要用于线性层权重和激活,以进一步减少显存和提高吞吐量。
学生需要记住:
- 精度越低:
- 显存越省;
- FLOPs 越多(单位时间内可执行操作更多);
- 但数值稳定性越差,需要更复杂的训练技巧。
6.2 如何“兼得二者”:混合精度训练
PPT 中提出核心问题:
“How can we get the best of both worlds?”
并给出结论:
- 解决方案:
- 默认使用 float32,但在可以的地方使用 bfloat16 / fp8。
更具体的计划:
-
在 前向传播 中,对 激活 (activations) 使用
bfloat16或fp8:- 这可以显著减少激活显存,占用约为原来的 1/2 或 1/4;
- 同时让 matmul 等算子利用 GPU 中对低精度的专用加速单元。
-
对于 参数 (parameters) 与 梯度 (gradients),仍使用
float32:- 保证累积更新过程中的数值稳定性;
- 减少因量化带来的损失。
-
利用深度学习框架的 自动混合精度 (Automatic Mixed Precision, AMP) 工具:
- PyTorch 官方文档:
- https://pytorch.org/docs/stable/amp.html
- NVIDIA 关于混合精度训练的指南:
- https://docs.nvidia.com/deeplearning/performance/mixed-precision-training/
- PyTorch 官方文档:
-
对于更激进的方案,可以使用 NVIDIA 的 Transformer Engine 支持 FP8:
- 将 FP8 更广泛应用于训练过程中的线性层;
- 参考 [Peng+ 2023] 等工作。
PPT 中再次强调:
- “How can we get the best of both worlds? Solution: use float32 by default, but use {bfloat16, fp8} when possible.”
- 这段内容重复出现,说明这是本讲非常核心的一条工程结论。
学生需要记住:
- 混合精度训练的核心思想:
- 把“对精度要求较高”的部分留给 float32;
- 把“计算量大且容忍误差”的部分交给 bfloat16 / fp8;
- 实际上,几乎所有现代大模型训练都会使用某种形式的混合精度,否则成本过高。
6.3 A100 与 H100 硬件参数对比表
PPT 中给出了一张关于 H100 SXM 与 H100 NVL 的对比表,其中列出了:
- FP32/FP16/BF16/TF32/INT8 等不同数据类型下的 teraFLOPS;
- 显存容量:
- H100 SXM:80GB;
- H100 NVL:94GB;
- 显存带宽:
- H100 SXM:3.35TB/s;
- H100 NVL:3.9TB/s;
- 其他多媒体相关单元,如 NVDEC、JPEG 解码单元数量 等。
表格中的关键结论:
- H100 相对于上一代 A100,在多种精度下都提供了更高的 teraFLOPS;
- 结合前面提到的 HPC 应用柱状图,可以看到:
- 在实际任务(如 FFT、基因测序)中,性能提升可达到 6X–7X;
- 这意味着:
- 同样规模的模型和数据集,在 H100 上训练所需时间会大幅度缩短;
- 反之,在相同训练时间内,H100 能支持更大的模型或更多数据。
学生需要记住:
- 在做训练计划和预算时:
- 不能只写“用 GPU 训练”,而要明确是 哪一代 GPU、哪种精度;
- 只有这样,前面的 FLOPs 估算公式才有实际意义。
6.4 混合精度在训练循环中的集成
结合第 5 章的训练循环,可以将混合精度训练集成进去:
- 使用 PyTorch AMP:
- 在前向和损失计算中使用
torch.cuda.amp.autocast(); - 在反向时使用
torch.cuda.amp.GradScaler进行梯度缩放,避免 underflow。
- 在前向和损失计算中使用
直观步骤:
- 在前向时:
with autocast(dtype=torch.bfloat16):包裹logits = model(x)与loss计算;
- 在反向时:
- 使用
scaler.scale(loss).backward(); - 再调用
scaler.step(optimizer)与scaler.update();
- 使用
- 其余如数据加载、日志记录不变。
学生需要记住:
- 混合精度训练的 API 使用并不复杂,但需要:
- 对数值不稳定问题(NaN、Inf)保持警惕;
- 结合学习率、梯度裁剪等手段一起调试。
小结:本讲应掌握的关键要点
-
资源核算思维:
- 训练时间和成本可以通过简单的 FLOPs 公式估算;
- 例如训练一个 70B 参数模型在 15T tokens 上,使用 1024 张 A100,大致需要 约 144 天;
- 这样的估算依赖于:
Total FLOPs ≈ 6 N params N tokens \text{Total FLOPs} \approx 6 N_\text{params} N_\text{tokens} Total FLOPs≈6NparamsNtokens
-
显存与参数规模:
- 使用 AdamW 并在 float32 下训练时,每个参数约需要 16 字节 的显存;
- 在 80GB 显存的 A100 上,粗略能承载 约 40B 参数(忽略激活)。
-
PyTorch 基础与训练循环:
- 张量是深度学习的统一表示,所有数据都以张量形式在 CPU/GPU 间移动;
- 训练循环包含:数据加载 → 前向 → 损失 → 反向 → 更新 → 监控;
- Autograd 为梯度计算提供自动求导,但也带来了激活显存开销。
-
FLOPs 与 MFU:
- 矩阵乘法 FLOPs 约为 O ( m k n ) O(mkn) O(mkn),大模型中的主要计算量都来自此类操作;
- 通过实测时间与理论 FLOPs,可以计算 MFU (Model FLOPs Utilization);
- MFU 超过 0.5 通常说明模型在硬件上利用得比较充分。
-
数值精度与混合精度训练:
- float32 提供稳定性但成本最高;bfloat16 和 fp8 带来大幅加速与显存节省;
- 混合精度训练 通过“前向激活用低精度、参数与梯度用高精度”为主流方案;
- PyTorch AMP 与 NVIDIA Transformer Engine 提供了实际工程中的落地工具。
-
硬件差异的重要性:
- A100 与 H100 在理论 FLOPs 和实际应用性能上差距巨大;
- PPT 中的图表展示了在 HPC 应用(3D FFT、基因测序)上可达 6X–7X 的性能提升;
- 做任何关于大模型训练的规划时,都必须清楚使用的是哪一代 GPU 与哪种精度。
通过本讲,读者应当不仅掌握 PyTorch 搭建模型和训练循环的基本技巧,更重要的是形成一种以显存和 FLOPs 为中心的工程思维:
- 在写下每一行代码前,先在心里大致算一算“这行代码要花多少钱”;
- 只有把资源账算清楚,才有可能在真实的大规模系统中高效地训练和部署大模型。
👉 更多笔记内容,关注公众号【你的生产力】免费领取


更多推荐


所有评论(0)