CS336课程笔记:lecture2 pytorch手把手搭建 LLM


目录

0. 本讲概览与学习目标

本讲是本课程的第二讲,延续上一节关于 语言模型 (Language Model)从零实现 的总览,开始真正动手用 PyTorch 搭建模型,并围绕资源效率(时间与显存)做系统的分析与估算。

核心问题:

  • 在给定 GPU 资源(如 A100/H100 数量、显存大小、FLOPs 能力)的情况下:
    • 能训练多大的模型?(参数规模)
    • 要花多长时间?(训练时长)
    • 如何在 PyTorch 中写出既正确又高效的代码?

解决方案路径:

  1. 步骤 A:内存与张量基础
    理解张量数据类型、维度与在 GPU 上的存储方式,学会按字节数估算显存占用,并从中得出“参数最多能有多少”、“激活会占多少”等结论。

  2. 步骤 B:计算量与 FLOPs 分析
    分析矩阵乘法、线性层等核心算子的 FLOPs,推导类似 6 N P 6NP 6NP 这种用于估算 LLM 训练开销的经验公式,并理解 MFU (Model FLOPs Utilization) 的含义。

  3. 步骤 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
  • 目标:估算 训练完一次的时间

解决思路:

  1. 先估算 总 FLOPs

    • 经验公式:
      Total FLOPs ≈ 6 × N params × N tokens \text{Total FLOPs} \approx 6 \times N_\text{params} \times N_\text{tokens} Total FLOPs6×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 统计,本讲后面会解释其来源。
  2. 再根据 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
  3. 用“需要的 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
  • 假设暂时不考虑激活占用,只关心 参数 + 梯度 + 优化器状态
  • 目标:估算 最多可训练的参数量

解决思路:

  1. 记住一个经验数:

    • 对 AdamW:
      • 参数本身:1 份;
      • 梯度:1 份;
      • 一阶矩动量 m m m:1 份;
      • 二阶矩动量 v v v:1 份;
    • 如果这些都用 float32 (4 bytes) 存储,则每个参数大约需要 16 字节
  2. 用显存总量除以每个参数所需的字节数:
    N params ≈ 80   GB 16   bytes N_\text{params} \approx \frac{80\,\text{GB}}{16\,\text{bytes}} Nparams16bytes80GB

  3. 粗略结果:

  • 得到的数量级约为 40B 参数

讲师特别强调:

  • 这个估算 还没有算上激活 (activations),而激活占用和 batch size、序列长度 强相关,在作业中会很关键;
  • 即便如此,这个粗略计算已经足以让我们在设计模型时有大致概念,不至于“随手一写就 OOM”。

学生需要记住:

  • 内存核算 是训练大模型的第一步;
  • 以 AdamW 为例,“16 bytes / 参数” 是一个非常重要的经验数,后面会在混合精度中看到如何降低这个数字。

1.3 本讲结构总览

讲师把本讲总结为三个主线:

  1. Memory accounting(内存核算)

    • 张量 (tensor) 基础 开始,理解数据类型和在 GPU 上的存储;
    • 学会估算参数、优化器状态、激活的显存消耗;
    • 为后续的 模型规模上限batch size 选择 提供依据。
  2. Compute accounting(计算核算)

    • 从简单的矩阵乘法开始,统计 FLOPs
    • 推到线性层、残差连接、注意力等操作的大致 FLOPs;
    • 连接到一开始的 6 N P 6NP 6NP 公式与 MFU 概念。
  3. PyTorch primitives + training loop(PyTorch 原语与训练循环)

    • 实际在 PyTorch 中写代码:
      • 张量在 GPU 上的创建与操作;
      • 使用 torch.nn.Module 组织模型;
      • 构造优化器与训练循环;
    • 结合 随机性、checkpoint、混合精度 等工程实践,构建可复现实验。

这一段的核心信息:

  • 本讲不是在教“怎么调参”,而是在教“如何把资源算清楚”;
  • 这种从底层算起的习惯,会贯穿后续所有关于大模型训练的内容。

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=1kdi)×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} ARm×k B ∈ R k × n B \in \mathbb{R}^{k \times n} BRk×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 k1 次加法;
    • 整体 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} FLOPslinear2BDinDout

其中:

  • 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:前向 + 反向

在训练中,每一步需要:

  1. 前向传播 (forward pass):计算模型输出与损失;
  2. 反向传播 (backward pass):计算梯度;
  3. 参数更新 (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 token3×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 FLOPs6×Nparams×Ntokens

讲师在这里强调:

  • 精确常数因子并不重要,重要的是 会算数量级
  • 在实际工程中,我们常以这种经验公式为起点,然后再通过 profile 工具 进一步精细分析。

学生需要记住:

  • 训练开销与参数量和 token 数成正比
  • 把 token 数翻倍,训练时间几乎也会翻倍;
  • 把模型参数翻倍,如果其他条件不变,训练时间也会近似翻倍。

3.4 用 PyTorch 实测 FLOPs 与时间:tensor_operations_flops()

在 PPT 中,讲师展示了一个实验流程,用 PyTorch 时间函数来验证计算:

  1. 定义一个矩阵乘法函数 time_matmul(x, w)

    • 输入张量 x 和权重矩阵 w 放在 GPU 上;
    • 使用 torch.cuda.synchronize() 确保计时时间准确;
    • 返回一次 matmul 的实际耗时 actual_time
  2. 预先用公式计算该 matmul 的理论 FLOPs:

    actual_num_flops = 2 m k n \text{actual\_num\_flops} = 2 m k n actual_num_flops=2mkn

  3. 用公式:

    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

  4. 对比 GPU 厂商给出的 理论峰值 FLOPs

    • 通过类似 get_promised_flop_per_sec(device, x.dtype) 的辅助函数查表;
    • 得到 promised_flop_per_sec
  5. 计算 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=pparamsidi(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()

讲师给出了一个典型的训练循环骨架:

  1. 迭代数据 loader:
    • 从数据集中取出一个 batch。
  2. 将数据移动到 device
    • x = x.to(device)y = y.to(device)
  3. 前向传播:
    • logits = model(x)
    • 计算损失 loss = criterion(logits, y)
  4. 反向传播:
    • optimizer.zero_grad()
    • loss.backward()
  5. 参数更新:
    • optimizer.step()
  6. 记录指标与日志:
    • 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()

讲师简要提到数据加载的几个要点:

  • 使用 DataLoaderDataset 抽象;
  • 合理设置 num_workerspin_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 param4(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

更具体的计划:

  1. 前向传播 中,对 激活 (activations) 使用 bfloat16fp8

    • 这可以显著减少激活显存,占用约为原来的 1/2 或 1/4;
    • 同时让 matmul 等算子利用 GPU 中对低精度的专用加速单元。
  2. 对于 参数 (parameters)梯度 (gradients),仍使用 float32

    • 保证累积更新过程中的数值稳定性;
    • 减少因量化带来的损失。
  3. 利用深度学习框架的 自动混合精度 (Automatic Mixed Precision, AMP) 工具:

    • PyTorch 官方文档:
      • https://pytorch.org/docs/stable/amp.html
    • NVIDIA 关于混合精度训练的指南:
      • https://docs.nvidia.com/deeplearning/performance/mixed-precision-training/
  4. 对于更激进的方案,可以使用 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。

直观步骤:

  1. 在前向时:
    • with autocast(dtype=torch.bfloat16): 包裹 logits = model(x)loss 计算;
  2. 在反向时:
    • 使用 scaler.scale(loss).backward()
    • 再调用 scaler.step(optimizer)scaler.update()
  3. 其余如数据加载、日志记录不变。

学生需要记住:

  • 混合精度训练的 API 使用并不复杂,但需要:
    • 对数值不稳定问题(NaN、Inf)保持警惕;
    • 结合学习率、梯度裁剪等手段一起调试。

小结:本讲应掌握的关键要点

  1. 资源核算思维

    • 训练时间和成本可以通过简单的 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 FLOPs6NparamsNtokens
  2. 显存与参数规模

    • 使用 AdamW 并在 float32 下训练时,每个参数约需要 16 字节 的显存;
    • 在 80GB 显存的 A100 上,粗略能承载 约 40B 参数(忽略激活)。
  3. PyTorch 基础与训练循环

    • 张量是深度学习的统一表示,所有数据都以张量形式在 CPU/GPU 间移动;
    • 训练循环包含:数据加载 → 前向 → 损失 → 反向 → 更新 → 监控;
    • Autograd 为梯度计算提供自动求导,但也带来了激活显存开销。
  4. FLOPs 与 MFU

    • 矩阵乘法 FLOPs 约为 O ( m k n ) O(mkn) O(mkn),大模型中的主要计算量都来自此类操作;
    • 通过实测时间与理论 FLOPs,可以计算 MFU (Model FLOPs Utilization)
    • MFU 超过 0.5 通常说明模型在硬件上利用得比较充分。
  5. 数值精度与混合精度训练

    • float32 提供稳定性但成本最高;bfloat16 和 fp8 带来大幅加速与显存节省;
    • 混合精度训练 通过“前向激活用低精度、参数与梯度用高精度”为主流方案;
    • PyTorch AMP 与 NVIDIA Transformer Engine 提供了实际工程中的落地工具。
  6. 硬件差异的重要性

    • A100 与 H100 在理论 FLOPs 和实际应用性能上差距巨大;
    • PPT 中的图表展示了在 HPC 应用(3D FFT、基因测序)上可达 6X–7X 的性能提升;
    • 做任何关于大模型训练的规划时,都必须清楚使用的是哪一代 GPU 与哪种精度。

通过本讲,读者应当不仅掌握 PyTorch 搭建模型和训练循环的基本技巧,更重要的是形成一种以显存和 FLOPs 为中心的工程思维

  • 在写下每一行代码前,先在心里大致算一算“这行代码要花多少钱”;
  • 只有把资源账算清楚,才有可能在真实的大规模系统中高效地训练和部署大模型。

👉 更多笔记内容,关注公众号【你的生产力】免费领取

Logo

中国智能体开发者社区,聚焦智能体与大模型开发,提供前沿资讯、实用工具链、开源项目及行业案例。通过技术沙龙、开发者大赛等活动,促进经验交流与协作,助力开发者快速构建创新智能应用。

更多推荐