本页目录

MLSys I · 训练系统与并行策略

对标:CMU 15-849 / Stanford CS229S / MLSys 会议论文 | 前置:gpu 线、par-01(并行模型、Amdahl/Gustafson)、ai-course(反向传播) 机器学习的算法你在 ai-course 学过,MLSys 讲的是让它在真实硬件上跑得起、跑得快——大模型训练是当今 HPC 最烧钱的战场。这一页讲训练系统的核心:计算图与自动微分怎么落地、显存为什么总不够(以及怎么省)、以及当一张卡装不下时的多卡并行策略(数据/张量/流水线并行)。这是理解"为什么训练 GPT 要几千张卡几个月"的系统视角。

1. 计算图与自动微分:框架底下是什么

反向模式 autodiff:计算图前向存中间值、反向链式累积梯度。

图 mlsys-01.3反向模式 autodiff:计算图前向存中间值、反向链式累积梯度。

PyTorch/JAX 的核心是自动微分(autodiff)——你写前向计算,框架自动算梯度(ai-course 的反向传播的工程实现)。机制是计算图:每个操作是节点、数据流是边。

关键系统洞察训练的显存大头是"为反向保存的中间激活值",不是参数本身——这直接引出下面的省显存技术。

2. 显存:永远不够的资源

训练显存四部分(参数/梯度/优化器状态/激活),16 字节/参数条形图。

图 mlsys-01.2训练显存四部分(参数/梯度/优化器状态/激活),16 字节/参数条形图。

训练时 GPU 显存装着四样东西,缺一不可:

  1. 模型参数(如 7B 模型 fp16 = 14 GB)。
  2. 梯度(同参数量级)。
  3. 优化器状态(Adam 混合精度要存 fp32 主权重 + 动量 + 方差,约 12 字节/参数——优化器状态常是最大头)。
  4. 激活值(前向的中间结果,为反向保存,随 batch 和序列长度涨)。

省显存的核心技术(让大模型能在有限显存上训):

这些都是"在固定显存里训更大模型"的工程手段——你在 comfy 课看到"显存不够跑不了大模型"的抱怨,MLSys 就是系统性地解决它。你的 4060 Ti 16GB 训练时也天天面对这个约束。

3. 分布式训练:一张卡装不下时

三种并行:数据(切batch) / 张量(切层内) / 流水线(切层) 对比。

图 mlsys-01.1三种并行:数据(切batch) / 张量(切层内) / 流水线(切层) 对比。

大模型一张卡装不下、或要加速,就要多卡并行。三种正交的并行策略(可组合):

策略 切什么 通信 适用
数据并行 切 batch(每卡完整模型、不同数据) 每步 all-reduce 同步梯度 模型装得下、要加速(最常用)
张量并行 切单个层(如把大矩阵乘按列切到多卡) 层内频繁通信(要求高速互联 NVLink) 单层太大装不下
流水线并行 切层(不同卡放不同层,像流水线) 层间传激活,要处理"气泡" 模型很深、跨节点

真实的大模型训练是"3D 并行"——数据 × 张量 × 流水线三种组合,加上 ZeRO 分片,才能把万亿参数塞进几千张卡高效训练。"怎么切模型和数据到几千张卡"是 MLSys 最核心的系统设计问题

4. 数据管线与训练效率

GPU 很贵,不能让它等数据。训练系统的另一半是喂数据:

5. 练习与要点

例 1(显存算账) 用 Adam 训一个 7B 参数模型(混合精度),按经验法则约 16 字节/参数(fp16 参数 2 + fp16 梯度 2 + fp32 主权重 4 + 动量 4 + 方差 4)估算 ≈ 112 GB,还没算激活——理解"为什么 7B 模型单张消费卡训不动",以及 ZeRO/量化为什么必要。

例 2(选并行策略) 模型能装进单卡但想用 8 卡加速 → 数据并行;单层大到装不下 → 张量并行;模型极深跨多节点 → 流水线。给三个场景各选策略。把并行三件套用到决策

例 3(检查点换显存) 梯度检查点省了激活显存但多了一次前向计算——估算"显存 vs 时间"的交换比,判断何时值得。perf 线的"用计算换存储"在训练里的实例\(\blacksquare\)


下一页:MLSys II——推理系统与算子优化:模型训好之后,怎么让它服务得快、省、稳。这是 vLLM、你天天用的 Claude/DeepSeek 背后的引擎。