Back to writing

/ CS336

[CS336-02] PyTorch and Resource Accounting

Notes for Stanford CS336 Spring 2025 lecture 2: PyTorch, FLOPs, Memory.

1 minCS336 · PyTorch · FLOPs · Memory

这篇是 Stanford CS336 Spring 2025 第 2 讲 Course Materials 的学习笔记。主线是:写模型之前,先学会算资源账。

1. 为什么要做 resource accounting#

语言模型训练看起来是在写 PyTorch,但真正约束训练规模的是资源:参数占多少显存、activation 占多少显存、一次 forward/backward 有多少 FLOPs、GPU 理论峰值和实际利用率差多少。

如果不会算资源账,就很难回答这些问题:

  • 这个模型能不能放进单卡?
  • batch size 能开多大?
  • bottleneck 是算力、显存,还是通信?
  • 实际跑出来的速度离硬件峰值有多远?

CS336 把这一讲放在前面,是因为后面所有 systems 话题都依赖这个基本能力。

2. 数值类型:精度也是系统设计#

课程从 float32、float16、bfloat16、fp8 讲起。它们不只是数学格式,也直接影响显存、带宽和 tensor core 利用率。

float32 精度高、动态范围正常,但每个数 4 bytes,训练大模型太贵。float16 每个数 2 bytes,吞吐高,但指数位少,容易 overflow 或 underflow。bfloat16 同样是 2 bytes,但指数位更接近 float32,因此训练稳定性更好,现代大模型训练中非常常见。fp8 更进一步压缩,但通常需要更复杂的 scaling 和校准策略。

系统视角下,低精度有两层收益:

  • 内存占用下降,同样显存能放更多参数和 activation;
  • 数据移动减少,更容易提高 arithmetic intensity。

3. Memory accounting:不只是参数#

很多初学者只算模型参数显存,例如 1B 参数用 bf16 存就是约 2GB。但训练时远不止参数。

一次训练通常还需要:

  • model parameters;
  • gradients;
  • optimizer states,比如 Adam 的一阶和二阶动量;
  • master weights;
  • activations,用于 backward;
  • 临时 buffer 和框架开销。

因此使用 Adam 训练时,每个参数可能对应十几 bytes 级别的状态。模型能不能训练,往往不是只看参数大小,而是看参数、梯度、优化器、activation 的总和。

这也解释了为什么 ZeRO、FSDP、activation checkpointing、sequence parallel 等方法会很重要:它们本质上都在重新分配或减少这些内存项。

4. Compute accounting:从线性层开始#

第 2 讲用线性模型和 Transformer 模块说明 FLOPs 怎么估计。

一个矩阵乘法 A @ B 的核心成本大致是:

2 * m * n * k FLOPs

其中乘法和加法各算一次操作。Transformer 里的主要计算也可以拆成多个矩阵乘法:QKV projection、attention output projection、MLP 的 up/down projection 等。

经验上,训练一次 token 的成本常用近似:

training FLOPs ≈ 6 * number_of_parameters * number_of_tokens

这不是精确公式,但足够用来做早期估算和 sanity check。

5. MFU:硬件峰值不等于实际速度#

课程引入了 Model FLOPs Utilization,也就是 MFU:

MFU = actual FLOP/s / theoretical peak FLOP/s

它描述模型训练实际吃到了多少硬件峰值。

MFU 低不一定是 PyTorch 写错了,也可能来自:

  • kernel launch overhead;
  • 小矩阵无法充分利用 tensor core;
  • memory bandwidth bottleneck;
  • activation 读写太多;
  • 通信和同步开销;
  • batch 或 sequence shape 不友好。

这也是后续 GPU、kernel、parallelism 课程的铺垫:训练速度不是只由 FLOP 数决定,数据移动和执行模型同样关键。

6. PyTorch 不是魔法#

这一讲还通过简单模型展示了 parameter initialization、tensor memory、pinned memory、CPU/GPU 数据传输等 PyTorch 细节。

PyTorch 的好处是抽象高,坏处是很多资源开销被隐藏。如果想写高性能训练系统,就需要在抽象之下继续追问:这个 tensor 在哪里?dtype 是什么?需要存到 backward 吗?会不会触发额外 copy?kernel 是否融合?

7. takeaway#

第 2 讲的核心不是某个 PyTorch API,而是一种习惯:写模型之前先算账。

  • 显存账:参数、梯度、优化器状态、activation 分别是多少;
  • 计算账:主要 FLOPs 来自哪些矩阵乘法;
  • 利用率账:理论 FLOPs 和实际吞吐之间差在哪里。

后面讨论 GPU kernel、分布式训练和推理优化时,这些账会不断出现。

参考#