这两周在消化 CMU 11-868 大语言模型系统课程的第五讲,主题是深度学习框架设计(Deep Learning Framework Design)。说实话,这门课之前几讲都在讲大模型训练、推理的系统架构,到这一讲突然落到“框架”这个层面,我当时第一反应是:都 2025 年了,我们天天用 PyTorch,还需要专门学框架设计吗?听完课之后我发现,恰恰是这些习以为常的东西,在大模型场景下成了一切瓶颈的源头。今天这期笔记就把我整理的内容和课后做的实验一起记录下来。
这门课面对的人群,不是“想用大模型写 demo”的开发者,而是那些需要踩进 LLM 底层、做分布式训练、做推理优化、甚至要自定义算子的工程师和研究者。深度学习框架是算法和硬件之间的那层“水泥”,它决定了你的梯度怎么流动、计算图怎么构建、显存怎么分配、算子怎么跑进 GPU。理解了框架设计,你再去看环形 All-Reduce、张量并行、FlashAttention、Zero 显存优化,不再是背公式,而是能看到它们各自长在框架的哪个环节。
1. 这门课为什么先从“框架设计”开始讲 LLM 系统
1.1 大模型时代的每一点性能,都是框架“挤”出来的
我在课程第一周的阅读材料里看到一句话:大语言模型的每一次 forward 和 backward,本质上都是在一个固定结构中重复执行成百上千个算子。你用 PyTorch 写一句output = model(input),背后会发生什么?自动微分建图、算子调度、显存分配、kernel launch、可能还有编译优化。深度学习框架就是负责把这些步骤组织起来,并且尽量让每一步都贴近硬件极限。
在普通模型的训练里,框架多花 10% 的开销可能不痛不痒;但在千亿参数、上万张 GPU 的训练任务里,一个算子调度多了几次 CPU 同步,一天下来就不知道浪费了多少 GPU 时。这就是为什么系统课会先把框架设计拎出来讲:框架设计的好坏,直接决定了大模型系统能在多大规模跑起来、能不能跑得又快又稳。
从我个人的学习体验来说,只有先搞清楚框架替我们做了什么,才能在遇到“显存爆了”“训练慢得像蜗牛”这类问题的时候,知道该去哪一层排查。以前的思维是“代码没错就行”,现在的思维是“代码没错只是开始,调度和内存访问才是大头”。
1.2 框架、IR 和编译器:它不是黑盒子
课程里把深度学习框架拆成了几个层级,我印象很深:最上层是用户 API,就是nn.Module、Optimizer这些东西;中间是计算图和中间表示 IR,负责描述计算逻辑;再往下是算子库和 kernel 实现;最底层是运行时 runtime,负责管理显存和调度执行。
这个分层和编译器非常像。LLVM 有前端、中间表示和后端,深度学习框架其实也有类似的“编译器”属性。PyTorch 早期更像一个“解释器”,动态地把算子一个个丢到 GPU 上;JAX 更激进,直接用 XLA 做整图编译;而 PyTorch 后来推出torch.compile,本质上也是在动态图外面包裹一层静态编译的壳,把 Python 代码先转换成 FX Graph 或者 TorchScript IR,再做融合和优化。
课程里给了一张对比图,说动态图适合研究和调试,静态图适合部署和极致性能。大语言模型训练规模那么大,完全动态图的开销是很高的,所以现在的主流做法是“让用户在 Python 里写,在 IR 层做优化”,也就是既要有动态图的灵活性,又要有静态图的性能。
1.3 我们从这门课里能带走什么
学完这一讲,我不再只是从“用户视角”看待框架。今后我们在设计一个训练脚本、写一个自定义算子、或者做推理优化时,都会先问自己几个问题:
- 计算图是动态建还是静态建?这一步有没有可能被编译器融合?
- 这个算子在框架层面的生命周期是什么?显存什么时候分配、什么时候释放?
- 反向传播的梯度图是不是也被框架隐式构造了?有没有额外开销?
- 如果我要做多机多卡,框架的并行抽象是否支持干净?我是否被 API 绑死?
带着这些视角去大量阅读源码和反汇编性能分析,比单纯调参有意思得多。
2. 框架设计最核心的三件事:图、算子、运行时
2.1 自动微分:动态图和静态图的恩怨情仇
自动微分是框架的“心脏”。课程重点讲了两条技术路线:反向模式自动微分(backpropagation)是训练神经网络的基本方法,但框架在实现时,可以选择先构建一个完整的数据流图,也可以选择“走一步看一步”。
PyTorch 的做法是动态图,也叫 define-by-run:你在 Python 里做一行张量操作,它就立刻构建一个节点,记录操作类型和输入输出关系,同时挂上对应的反向函数。这种方式的优点是调试方便,可以用原生的 Python 控制流,但代价是每一次训练迭代,计算图都要重新构建,而且 Python 和 C++ 之间频繁交互的开销很大。
TensorFlow(老版本)和 JAX 更像静态图思路。JAX 里的jax.grad会先对你的函数做转换,得到一个新的函数,再对函数体做编译优化,生成高效的 XLA HLO IR。静态图的好处是编译器能看完整张图,可以做算子融合、内存规划甚至并行调度。代价是写代码时要适应函数式编程约束,不能用随意的 Python 动态控制流。
课程里举了一个非常直观的例子:写一个带for循环的网络,动态图可以直接在 Python 里写for,静态图则需要像jax.lax.scan这样的结构来让编译器“看到”循环次数。为了兼得两者,PyTorch 2.x 的torch.compile用 TorchDynamo 在 Python 字节码层面“拦截”张量操作,把它们拼接成一张图,再交给后端编译器优化。这算是两种范式的融合。
2.2 算子融合:为什么 FlashAttention 能“封神”
在大语言模型里,注意力机制的计算量占了很大一部分。如果用朴素方式实现,每个头的 QK^T、softmax、attention @ V 会产生大量的中间张量,比如QK^T的矩阵结果和 softmax 后的概率矩阵。这些中间张量要写回显存,下一次计算再读出来。对于几十亿参数的模型,序列长度一长,这些中间张量能占据上百 GB 显存。
FlashAttention 的核心思路说起来并不复杂:算子融合 + 分块计算。它把前面几个计算步骤融合成一个 kernel,在 GPU 的 SRAM 中完成 query、key、value 的分块读取和计算,只把 final output 和统计量写回 HBM。框架设计在这个场景中扮演的是“能否支持和暴露融合操作”的角色。如果你用 PyTorch 直接写多头注意力的各个步骤,它会老老实实地每一步落一次显存;但如果你调用torch.nn.functional.scaled_dot_product_attention,在支持的后端上就可以直接触发融合的 FlashAttention kernel。
课程里还提到,算子融合不是只有注意力。比如 MLP 里的Linear+ReLU+Dropout,在推理时可以融合成一个 kernel,避免中间结果的多次读写。框架要做的事,是提供一个高层 API,然后在底层判断硬件、形状、数据类型,决定是走融合 kernel 还是走复合算子。这就是“图优化”阶段的核心工作之一。
2.3 并行策略:框架如何“接管”多卡协同
大模型训练绕不开数据并行、模型并行、流水线并行这些概念。课程里专门强调,现代框架在设计 API 时必须把这些并行策略变成“声明式”的,而不是让用户手动做通信。
PyTorch 的DistributedDataParallel本质上是在反向传播时对梯度做 AllReduce,用户几乎感觉不到。FSDP(Fully Sharded Data Parallel)更进一步,把参数、梯度和优化器状态分片到所有 GPU 上,在需要时再 all-gather 组装。这些功能如果没有框架层支持,单靠用户在纯 Python 里手搓通信代码,很容易出现同步错误和死锁。
从框架设计的角度看,它在职责上做了一个漂亮的划分:用户定义模型和 forward/backward,框架自动插入通信原语。并且,框架可以通过计算图分析在什么位置插入通信最合适,比如在反向传播过程中梯度产生的第一时间就启动 AllReduce,让通信和计算重叠。这就是为什么同一套多卡训练,使用框架内置策略比自己写分布式代码高效得多。
这个部分让我意识到:好的框架设计是在正确的位置提供抽象,而不是给你一堆 API。抽象得越高,普通开发者越容易上手,但也意味着框架需要做更多自动化决策。
3. 课后实验:从零手写一个 mini 自动微分框架
3.1 为什么我要做这个“造轮子”实验
课程讲到自动微分时,我总觉得“反向传播不就是链式法则嘛”这种理解太浅了。为了真正理解动态图框架的内部流程,我花了两个晚上写了一个极简自动微分框架,只依赖 NumPy。目标不是复刻 PyTorch,而是搞清楚下面几个关键问题:
- 反向传播的图遍历到底怎么设计?
- 每个算子如何注册自己的反向函数?
- 梯度累积到叶子节点时,为什么是“累加”而不是“覆盖”?
- 计算图的拓扑排序为什么是反向传播正确性的关键?
写完这个小框架后,再回头看 PyTorchTensor里的grad_fn、backward、retain_graph这些属性,突然就全部串起来了。
3.2 核心数据结构:从一个带梯度的张量开始
我实现的核心是一个带grad和_backward的 Tensor 类。data保存数值,grad保存梯度,_backward是该节点到父节点的“梯度传播函数”。每一个操作返回的新 Tensor 都会记录它的父节点(children),并用闭包保存反向逻辑。
这里有个关键设计:反向函数不是立刻执行的,而是在backward()被调用后,按照拓扑序从输出到输入逐个触发。我把_prev这个集合保存成一个节点的前驱,方便后续做拓扑排序。
简单的结构如下:
import numpy as np class Tensor: def __init__(self, data, requires_grad=False, children=(), op=''): self.data = np.array(data, dtype=np.float32) self.requires_grad = requires_grad self.grad = None self._backward = lambda: None self._prev = set(children) self._op = op这里有几个细节需要注意:
data统一转成 NumPy 数组,保证矩阵运算的接口一致。requires_grad表示该张量是否参与梯度计算,类似 PyTorch 里的叶子节点控制。_backward是节点自己身上挂的反向传播逻辑,由运算符创建时定义。_prev是为了构造计算图,等到backward()时能沿着_prev做拓扑排序。
3.3 反向传播的拓扑排序实现
反向传播时,我先从当前张量出发,通过_prev深度优先遍历整张计算图,得到一个拓扑序列。然后初始化输出节点的梯度为 1(标量输出时),再逆拓扑序逐个调用每个节点的_backward()。
这个逆序很关键:它确保了在计算某个节点的梯度时,它所有子节点的梯度都已经计算完成,符合反向传播的“自顶向下”依赖关系。
def backward(self): topo = [] visited = set() def build_topological_order(v): if v not in visited: visited.add(v) for child in v._prev: build_topological_order(child) topo.append(v) build_topological_order(self) self.grad = np.ones_like(self.data) for node in reversed(topo): node._backward()如果不做拓扑排序,而是一看到节点就立即算反传,很可能会漏算依赖,或者访问到尚未计算梯度的节点。这也是动态图框架里backward()必须遍历图的原因——计算图在运行完前向后,本身就是一个无环有向图(DAG)。
3.4 常用算子的反向函数注册
接下来实现几个基础算子。以加法为例,z = x + y,对 x 的梯度就是 z 的梯度,对 y 的梯度也是 z 的梯度。如果 x 和 y 同时需要梯度,两边要做“累加”。因为一个变量如果参与了多个算子,它会有多条梯度路径,最终梯度是这些路径的和。
def add(a, b): out = Tensor(a.data + b.data, requires_grad=a.requires_grad or b.requires_grad, children=(a, b), op='+') def _backward(): if a.requires_grad: a.grad = (a.grad if a.grad is not None else 0) + out.grad if b.requires_grad: b.grad = (b.grad if b.grad is not None else 0) + out.grad out._backward = _backward return out矩阵乘法z = a @ b的反向要稍微绕一点:对 a 的梯度是out.grad @ b.T,对 b 的梯度是a.T @ out.grad。维度对不上的时候,多枚多次就会搞混,所以务必拿笔验算一下。
def matmul(a, b): out = Tensor(a.data @ b.data, requires_grad=a.requires_grad or b.requires_grad, children=(a, b), op='@') def _backward(): if a.requires_grad: a.grad = (a.grad if a.grad is not None else 0) + out.grad @ b.data.T if b.requires_grad: b.grad = (b.grad if b.grad is not None else 0) + a.data.T @ out.grad out._backward = _backward return outReLU 的反向更简单:前向时大于零的位置保留梯度,小于等于零的位置梯度置零。
def relu(a): out = Tensor(np.maximum(a.data, 0), requires_grad=a.requires_grad, children=(a,), op='relu') def _backward(): if a.requires_grad: a.grad = (a.grad if a.grad is not None else 0) + out.grad * (a.data > 0) out._backward = _backward return out3.5 用 mini 框架训练一个两层网络
为了让实验更贴近真实场景,我用这个小框架搭建了一个两层 MLP,在随机数据上做回归任务,loss 函数是均方误差。这个例子虽然玩具,但能把算子、拓扑排序、梯度累积这些概念串起来。
np.random.seed(42) # 随即生成一份线性可分的数据 N = 64 in_features = 3 hidden_features = 8 out_features = 1 x = np.random.randn(N, in_features) w1_true = np.random.randn(in_features, hidden_features) w2_true = np.random.randn(hidden_features, out_features) y = np.sin(x @ w1_true) @ w2_true + 0.01 * np.random.randn(N, out_features)定义模型参数,并设置requires_grad=True:
W1 = Tensor(np.random.randn(in_features, hidden_features) * 0.01, requires_grad=True) b1 = Tensor(np.zeros((1, hidden_features)), requires_grad=True) W2 = Tensor(np.random.randn(hidden_features, out_features) * 0.01, requires_grad=True) b2 = Tensor(np.zeros((1, out_features)), requires_grad=True)前向传播看起来和 PyTorch 很像:
def forward(x_tensor): z1 = add(matmul(x_tensor, W1), b1) a1 = relu(z1) z2 = add(matmul(a1, W2), b2) return z2 x_tensor = Tensor(x) pred = forward(x_tensor)均方误差损失的反向需要特别注意标量输出时的梯度传导。我的实现是:
def mse_loss(pred, target): diff = pred.data - target loss = np.mean(diff ** 2) loss_tensor = Tensor(loss, requires_grad=pred.requires_grad, children=(pred,), op='mse') def _backward(): if pred.requires_grad: g = out.grad # scalar pred.grad = (pred.grad if pred.grad is not None else 0) + (2.0 * diff / diff.size) loss_tensor._backward = _backward return loss_tensor然后开始训练:
learning_rate = 0.01 for step in range(200): pred = forward(x_tensor) loss = mse_loss(pred, y) for p in [W1, b1, W2, b2]: p.grad = None loss.backward() for p in [W1, b1, W2, b2]: p.data -= learning_rate * p.grad if step % 40 == 0: print(f"step {step}, loss = {loss.data:.6f}")我在测试时,loss 从最初的 1.2 左右稳步降到 0.01 附近,说明反向梯度生成了正确的方向。这个实验给我最大的启发是:框架里看似神奇的backward(),本质上就是一次拓扑序下的链式法则展开。每次调用backward(),底层都在做类似上面的遍历。
4. 框架设计如何影响大语言模型的训练与推理性能
4.1 显存管理:谁在偷偷吃掉你的 memory
大语言模型训练中,显存占用主要来自四块:参数、梯度、优化器状态、中间激活值。前三个和模型规模强相关,但激活值是由框架的计算图和算子实现决定的。如果框架对中间张量的生命周期管理不够好,激活值就会在显存里囤积。
PyTorch 有一个 caching allocator,它不会每次分配张量都向 CUDA 申请内存,而是维护一个缓存池,把释放的内存块留作下次重用。这样可以大幅度减少cudaMalloc的次数,毕竟cudaMalloc是很慢的系统调用。但是,缓存池也可能导致显存永远不会完全释放,即使你的脚本临时创建了一个大张量再删掉,显存占用可能依然居高不下。在课程演示里,教授用torch.cuda.memory_summary()检查一个训练进程,发现实际上有大量碎片化的缓存块,它们无法被复用,导致有效显存减少。
这也是为什么框架设计里会有“内存规划器”这种角色。静态图编译器可以在编译期分析每个中间张量的生存期,提前规划内存复用:如果两个张量的生命周期不重叠,就让它们共用同一块显存。而动态图很难做这样的全局规划,所以 PyTorch 采用了缓存池这种相对保守的策略。大模型训练中,激活值占用的显存极大,因此出现了激活重计算(gradient checkpointing)这种以算换显的技术。它本质上也是让框架“丢弃”中间的激活结果,在反向传播时重新计算,从而缩短中间张量的存活时间。
4.2 从 Dynamo 到 Inductor:PyTorch 是如何变快的
课程里用了挺大的篇幅讲 PyTorch 2.x 的编译路径。TorchDynamo 在 Python 字节码层面拦截你写的forward函数,把张量操作“扣”出来,转换成一种称为 FX Graph 的计算图。接着,Inductor 后端会拿到这张图,继续做算子融合和代码生成,最终生成 Triton kernel 或 CUDA kernel。
为什么这样做能让大模型变快?我举一个实际例子。Transformer 的 MLP 层里经常出现这样的序列:
linear(x) -> gelu -> dropout -> linear
在朴素动态图模式下,框架会分别执行四个 kernel,每一个 kernel 启动都有 CPU 端到 GPU 端的 launch 延迟,以及中间张量的显存读写。Inductor 可以把linear -> gelu -> dropout融合成一个 Triton kernel,这样所需时间大幅下降。实测显示,融合后的执行时间可能只有原来的三分之一到五分之一。
框架设计在这里做了一个关键取舍:Python 层的nn.Module只是“示意图”,真正执行的是编译后的融合 kernel。这种理念让大模型训练的重启成本变高,但稳态运行的吞吐更高。这个思路也被 JAX 和 TensorFlow 采用多年,现在大家殊途同归。
4.3 通信与计算重叠:隐藏在数据并行背后的框架魔法
大模型的多卡训练,通信量非常大。以数据并行为例,反向传播时需要同步所有 GPU 上的梯度,做一次 AllReduce。如果等所有层的梯度都算完再通信,GPU 在等待通信期间是空闲的。聪明的框架设计会用“梯度桶(gradient bucket)”:把参数分成若干桶,反向传播算完一个桶的梯度,立刻启动这个桶的 AllReduce,算下一个桶的同时通信在后台进行。
这让通信和计算尽可能重叠,训练吞吐能提升 20% 以上。在 FSDP 里,分片的参数可以在前向传播时按层 all-gather,用完后立即释放;反向传播时再次 all-gather 参与计算,用完后释放。这些调度逻辑如果交给用户手写,很容易出错;而框架可以把它们作为通信原语隐藏到 API 后面。
我想强调的是,这些优化不是“魔法”,它们全部基于框架对计算图的分析和运行时的调度。所以当你的大规模训练遇到性能瓶颈时,应该去检查框架是否生成了最优的计算图、通信是否重叠、kernel 是否融合,而不是单纯怀疑“是不是代码写得不够好”。
5. 常见问题与踩坑实录
5.1 “梯度对不上”怎么办?用数值梯度检查法
自己写自动微分时,最常见的 bug 就是反向传播公式写错,或者维度没有对上。我在实验里就出现过 3x3 矩阵乘法的梯度维度反了,导致参数在训练中直接 NaN。
一个非常有效的排查手段是数值梯度检查:对每个参数加一个微小扰动,用中心差分近似梯度,然后跟解析梯度做对比。如果误差在 1e-6 量级,说明反传公式基本正确。
def numerical_gradient(fn, x, epsilon=1e-6): grad = np.zeros_like(x) it = np.nditer(x, flags=['multi_index']) while not it.finished: idx = it.multi_index old_val = x[idx] x[idx] = old_val + epsilon f_right = fn(x) x[idx] = old_val - epsilon f_left = fn(x) x[idx] = old_val grad[idx] = (f_right - f_left) / (2 * epsilon) it.iternext() return grad检查时,注意这个函数接收的x是 NumPy 数组,而解析梯度可能需要转成 Tensor。数值梯度的误差来源包括浮点精度和模型非线性强度,如果误差在 1e-4 以内基本没问题。
5.2 显存爆炸时,怎么定位是谁在占用
大模型训练最痛苦的问题之一是显存溢出。很多人只会盯着nvidia-smi看总占用量,但无法判断是参数、激活值、还是中间张量占的。
PyTorch 提供了不错的排查工具。在代码里加上这几行:
import torch torch.cuda.reset_peak_memory_stats() # ... 运行前向/反向 ... print(torch.cuda.memory_summary())memory_summary会显示当前分配、峰值分配、缓存池大小和碎片情况。如果你想看每个操作的显存占用,可以用torch.profiler:
from torch.profiler import profile, ProfilerActivity with profile(activities=[ProfilerActivity.CUDA]) as prof: loss.backward() print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=30))profile 可以看到每个算子的 CUDA 时间,但显存占用需要结合torch.cuda.memory._record_memory_history()这一整套工具,才能查看每个张量分配栈。对于大模型训练,建议打开PyTorch nightly的 memory snapshot 功能,生成 HTML 快照,能非常直观地看到哪一行创建的大张量一直没释放。
5.3 自定义算子反而更慢,原因多半是“调度开销”
我写过不少自定义的torch.autograd.Function,比如一个融合的 FP8 量化算子。在单独 benchmark 时,确实比逐个调用 PyTorch 原生算子快;但放到大模型训练脚本里,整体速度可能反而变慢。原因是 Python 和 C++ 的边界调用是有固定开销的,如果你的自定义算子在每个 iteration 中只做很小的计算,调度开销会盖过计算收益。
一个更隐蔽的问题是:自定义算子可能阻止了框架的图优化。比如你在自定义 forward 里随机写一点 Python 逻辑(比如if condition: ...),TorchDynamo 可能无法将这个算子安全地并入图,于是整个forward都退化成 eager 模式。解决思路是尽量让自定义算子保持“纯函数”风格,并在没有特殊动态控制流时给算子写 torch.library/torch.ops 的类型声明,帮助编译器理解它的输入和输出。
课程里教授也提供了一个经验:如果要自定义算子,先确认它能被torch.compile捕获。一个简单的验证方式,是在模型上跑一次compiled_model = torch.compile(model),看一下编译日志里有没有fallback或graph break。如果 graph break 很多,说明你的代码结构阻碍了编译器视野,这时候就应该重构代码。
6. 课程之外的体会:框架设计是一种“系统思维”
6.1 框架选型或设计时,先回到计算图
走完这一讲,我最大的收获不是记住了某个 API,而是学会了一种“框架思维”。拿到一个模型或系统设计任务时,先画图:张量从哪来、经过哪些算子、中间结果需要存活多久、哪些操作之间没有数据依赖、哪些地方可以并行或融合。
这个思维放到大语言模型场景特别有用。比如推理服务里,如果我要做连续批处理(continuous batching),最核心的问题不是 Python 怎么调度请求,而是 GPU 上的显存如何按需分配给不同请求的 KV cache。这本质上就是一个动态显存分配器的设计问题,和 PyTorch caching allocator 要解决的问题一模一样。看懂框架设计的人,做起推理系统来会更有底气。
6.2 如果你也想深入自学框架设计
如果你也想提升这块能力,我建议按这个顺序尝试:
- 先用 PyTorch 手写一个
nn.Module,配置torch.compile,观察编译和未编译的性能差异。 - 结合
torch.profiler找出时间占比最高的算子,思考为什么它是瓶颈。 - 尝试用 Triton 写一个简单的融合 kernel,对比原生实现。
- 读一遍 micrograd 或 tinygrad 源码,理解自动微分如何用几十行代码实现。
- 有条件的话,去看 PyTorch 的
aten/src/ATen和torchinductor目录,了解算子注册和代码生成的关系。
这样一轮下来,你会发现自己看论文里各种系统优化方案时,不再是雾里看花。
我个人在写 mini autograd 时踩过的最大一个坑,是没有给叶子节点在每轮迭代前清空grad。如果不清空,多个 step 的梯度会累加,导致 loss 震荡甚至发散。这个细节在 PyTorch 里是通过optimizer.zero_grad()帮你处理好了,但自己在设计框架时就要考虑到这种生命周期管理。框架设计就是无数个这样的小决策堆出来的:每个决策单独看都不复杂,组合在一起就决定了系统的上限。
这次的《深度学习框架设计》笔记就写到这里。下一讲的内容会进入大语言模型分布式训练的具体策略,我准备把这次学到的图优化、显存管理与后面要讲的流水线并行对照着再看,相信到时候还会有新的收获。