资讯动态

深度学习框架设计核心:自动微分、算子融合与显存管理如何决定大模型性能

发布时间:2026/9/15 23:01:09 来源:尧图企业网站定制
这两周在消化 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 里的LinearReLUDropout在推理时可以融合成一个 kernel避免中间结果的多次读写。框架要做的事是提供一个高层 API然后在底层判断硬件、形状、数据类型决定是走融合 kernel 还是走复合算子。这就是“图优化”阶段的核心工作之一。2.3 并行策略框架如何“接管”多卡协同大模型训练绕不开数据并行、模型并行、流水线并行这些概念。课程里专门强调现代框架在设计 API 时必须把这些并行策略变成“声明式”的而不是让用户手动做通信。PyTorch 的DistributedDataParallel本质上是在反向传播时对梯度做 AllReduce用户几乎感觉不到。FSDPFully 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_gradFalse, children(), op): self.data np.array(data, dtypenp.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_grada.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_grada.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_grada.requires_grad, children(a,), oprelu) 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_gradTrueW1 Tensor(np.random.randn(in_features, hidden_features) * 0.01, requires_gradTrue) b1 Tensor(np.zeros((1, hidden_features)), requires_gradTrue) W2 Tensor(np.random.randn(hidden_features, out_features) * 0.01, requires_gradTrue) b2 Tensor(np.zeros((1, out_features)), requires_gradTrue)前向传播看起来和 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_gradpred.requires_grad, children(pred,), opmse) 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(fstep {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 到 InductorPyTorch 是如何变快的课程里用了挺大的篇幅讲 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, epsilon1e-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.profilerfrom torch.profiler import profile, ProfilerActivity with profile(activities[ProfilerActivity.CUDA]) as prof: loss.backward() print(prof.key_averages().table(sort_bycuda_time_total, row_limit30))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()帮你处理好了但自己在设计框架时就要考虑到这种生命周期管理。框架设计就是无数个这样的小决策堆出来的每个决策单独看都不复杂组合在一起就决定了系统的上限。这次的《深度学习框架设计》笔记就写到这里。下一讲的内容会进入大语言模型分布式训练的具体策略我准备把这次学到的图优化、显存管理与后面要讲的流水线并行对照着再看相信到时候还会有新的收获。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价