PyTorch大模型训练显存核算:从OOM到精准预算
CS336 的前几次课其实是很多做LLM训练的人最该补的基础课。很多人能把模型跑起来却说不清楚自己的显存到底花在了哪里遇到OOM第一反应是调小batch size调完还是炸然后开始乱试。我第一次在40GB单卡上训练一个1.3B参数模型时也这样模型文件才2.6GB直觉上怎么都不至于把整块卡吃满结果一个step下来直接OOM。后来认真看CS336的课程材料才发现这门课在很早的阶段就要你回答一个问题在把模型交给PyTorch之前你能不能先手算出一次forward/backward要消耗多少显存、多少算力。这节课的名称就叫“PyTorch, resource accounting”既讲PyTorch基础框架的用法又讲资源的“算账”方法。这篇文章我会把这条主线拆开讲为什么资源核算那么重要、PyTorch里哪些底层设计直接影响资源开销、怎么从零手动估算训练一个Transformer的显存、以及用哪些工具真正把账算明白。适合准备认真刷CS336作业的读者也适合任何打算自己训练或微调语言模型、但总觉得显存瓶颈说不清的工程师。1. 为什么CS336在入门阶段就把资源核算单独立了一讲先交代一下背景。CS336这轮课名字叫“Language Modeling from Scratch”意思是所有组件都不直接拿现成的库糊上去数据加载你要自己写模型结构要用PyTorch拼优化器、混合精度、分布式通信都要自己实现一遍。很多第一次接触这门课的人会有一个错觉觉得最大的难点是模型架构本身。但真正动手做作业的时候会发现卡住你的往往不是Transformer的attention怎么写而是“这个配置一下去显存到底够不够”“这个batch size到底能不能跑”“多卡并行的时候每张卡的消耗是不是均匀的”。resource accounting资源核算解决的就是这一类问题。它的核心思想非常朴素在训练一个模型之前先把必要的账目算清楚而不是把任务丢给GPU去“试错”。算清楚之后很多事情就变得非常直接了——该不该上梯度累积、要不要开gradient checkpointing、序列长度能开到多少、用几路数据并行这些配置在你真正启动训练前就已经有结论了。为什么这门课要在这么早的阶段专门花一整节讲它因为后续的作业几乎每一环都在依赖这个能力。你写了一个自定义的分布式通信逻辑怎么知道它有没有带来额外的显存开销你手写了一个混合精度优化器怎么验证它确实比原来省了资源如果对资源消耗没有“量”的直觉这些作业做起来就是盲人摸象只能靠不断看报错来试。这里给一个非常直观的对比。一个1.3B参数的语言模型用常见的bf16混合精度训练在不考虑激活内存的情况下模型参数、梯度、优化器状态加起来就需要大约20GB左右的显存。很多初学者拿到模型文件只有2.6GB以为几块卡随便跑结果一个micro batch塞进去就显存溢出。反过来如果你提前算过账就会知道这类模型哪怕单卡跑也至少要一块比较充裕的显存卡然后设计好micro batch和梯度累积方案整个过程可以少踩很多坑。资源核算说起来抽象其实就三步知道自己有多少资源、知道自己训练一个批次要花多少资源、知道瓶颈在哪个环节。CS336这节的关键价值就是帮你在训练代码运行之前建立这三方面的认知。2. PyTorch基础框架中决定资源开销的四个设计细节PyTorch不是黑盒它对显存的控制方式其实是有迹可循的。这里我不打算推一遍API文档而是挑四个直接决定资源开销的设计细节把它们讲透。这四个点不搞明白后面的账就算不准。2.1 Tensor的存储与视图共享PyTorch里的Tensor存储和视图是分开的。简单说一个Tensor背后可能有一块独立的显存也可能只是另一块显存的一个“视图”。当你做切片、转置、expand这类操作时新得到的Tensor和原Tensor共享同一块底层存储不会额外占用资源。但一旦你调用了contiguous()、clone()、copy_()这类会改变数据排列或复制数据的操作就会产生一份新的显存分配。这个细节在实际训练里非常容易踩坑。比如你在预处理batch时为了把不同长度的序列拼成矩阵做了一次pad然后出于习惯调用了.contiguous()这一步可能就让输入数据多占了一份显存。再比如你在某个模块的forward里对激活值做了切片并且后面又做了view看起来只是角度变化但某些算子底层需要连续内存瞬间又会触发一次复制。做资源核算的时候如果不理解存储和视图的区别你会经常遇到“显存比预期的多出好几个GB”的困惑因为很多开销根本不是参数带来的而是数据在搬运和重排过程中产生的临时分配。有一个很实用的判断原则不需要改变数值、只改变形状或观察角度的操作一般共享存储需要把内存排布变成连续、或者生成一份独立副本的操作一定会产生新分配。遇到可疑操作直接用x.storage().data_ptr()看看地址是否一致比猜来猜去快得多。2.2 autograd计算图与激活的存与舍第二个关键设计是autograd。训练模式下PyTorch在forward过程中会把参与求导的中间结果保存下来等backward时再拿出来用。这些中间结果就是训练内存的重要来源通常被称为激活内存activation memory。你以为显存主要是模型参数实际很多场景下激活内存才是那个偷偷吃显存的大头。这里有个值得记住的对比inference和training同样的模型、同样的batch显存消耗可能相差好几倍。原因是inference阶段有torch.no_grad()包裹PyTorch不需要构建计算图中间结果算完就可以丢掉而训练阶段每一层的输出都要留着给反向传播用。所以你想评估一个大模型到底多占显存光用模型参数的体积来估算一定会严重低估。这就引出一些实用策略不参与训练的层随手设置requires_grad_(False)它就不会保存梯度相关的中间结果模型里某些自定义算子如果自身不需要梯度也可以用torch.no_grad()包起来减少计算图节点。CS336作业里有一个常见要求是手动实现一个简化版Transformer很多同学在调试benchmark时发现显存增长异常最后定位到的问题往往就是没有管理好哪些张量进入了计算图。2.3 device与数据搬运的隐藏成本第三个细节是CPU与GPU之间的数据搬运。PyTorch里Tensor从CPU搬到GPU或者从GPU搬回来都是一个相对昂贵的操作。更隐蔽的是它会打断GPU的异步流水线造成同步等待进而影响整体吞吐。很多人在做资源核算时只盯着显存忽略了数据传输对耗时的影响。DataLoader里面有个参数叫pin_memory很多人把它当成默认不动的东西。实际上pin_memoryTrue会让DataLoader使用锁页内存来临时存放数据这样从CPU复制到GPU的速度会明显变快。但注意锁页内存属于系统内存而不是显存它不会增加你的显存占用却会显著增加宿主机的内存消耗。如果你的服务器本身内存不多把num_workers拉高再加上pin_memoryTrue可能模型还没开始训练CPU内存先告急了。这一类隐藏成本在资源核算里也要算进去。CS336作业在实现数据流水线时会鼓励你自己去测量数据加载部分的耗时。如果你用nvidia-smi只盯着显存看完全忽略Dataloader对CPU内存和GPU流水线的影响测出来的训练速度会严重失真。2.4 混合精度下的双轨存储混合精度是现在训练大模型的默认配置但它带来的资源变化不是简单的“省一半”。用bf16做前向和反向模型会维护一个bf16的实时权重副本但优化器更新时通常还需要一份fp32的master weight和Adam状态。所以混合精度训练下参数在显存里是“双轨”存在的一份低精度用于计算一份高精度用于更新。很多初学者第一次打开优化器代码发现里面维护的不只一组参数往往很困惑。其实你只需要记住一个结论bf16混合精度相比纯fp32训练主要省的是梯度和计算时的内存带宽但并没有把优化器状态直接减半因为Adam这类优化器内部的一阶矩、二阶矩仍然是fp32精度。到了后面做作业你自己去实现mixed precision trainer的时候会亲手把这两条存储轨道的账算一遍这个理解才算真正落地。3. 从参数开始手动推导一次训练资源账目现在到了最核心的部分动手算账。我们先忽略一些运行时细节从最基本的公式出发把一次训练的资源消耗估算出来。3.1 静态内存四件套训练一个模型显存里至少存在四类固定的东西模型参数、梯度、优化器状态、以及运行时上下文。前三类可以按参数数量N直接估算它们是静态的不随batch大小变化。假设你有一个N个参数的网络模型参数如果以bf16存储每参数2字节梯度在混合精度训练中通常也以bf16保存每参数2字节优化器状态以Adam为例一般会保存fp32的master weight、一阶矩m、二阶矩v每一项都是4字节共12字节。所以一个粗略的快速估算方法是bf16混合精度 Adam训练静态显存成本大约为每参数 2 2 12 16字节。结合我们前面的说法2是bf16权重2是bf16梯度12是优化器相关的三份fp32状态。这个“16字节/参数”的经验公式在行业里很常用适用于快速判断能不能用单卡跑。举个例子一个7B参数的模型静态成本大约是7e9 × 16字节 ≈ 112GB。这意味着哪怕你用梯度累积把batch size压到很小光模型、梯度、优化器状态就要占一百多GB显存。这也就是为什么7B级别的预训练通常要上多卡甚至CPU offload不是没有道理的。相比之下如果只是做推理不需要梯度和优化器状态bf16推理静态成本就是每参数2字节7B模型约14GB差距一目了然。这里要提醒一点不同框架对混合精度的实现细节有差异。有的框架会在优化器状态里额外给梯度也保留fp32版本这时每参数的静态成本会到18字节左右。所以16字节是一个“至少成本”实际落地时建议以你使用的框架为准不要拿一个公式通吃所有情况。3.2 激活内存的大头从哪来静态成本算清楚了接下来是动态部分激活内存。激活内存和batch size、序列长度、隐藏维度、层数直接相关与参数量关系不大。这也是为什么有时候一个模型很大激活却不吓人而一个模型不大seq_len一拉长显存照样爆掉。一个比较常用的粗略公式是激活显存约等于 batch_size × seq_len × hidden_size × (34 5 × num_layers) 字节。这里的34和5是经验系数来自Transformer前向过程中需要保存的各类中间张量的累加对你理解量级已经足够。代入一个1.3B规模的例子假设hidden_size2048num_layers24seq_len1024micro batch为4。激活量大概是 4 × 1024 × 2048 × (34 5 × 24) ≈ 4 × 1024 × 2048 × 154 ≈ 1.29GB。这个数字不算大。但如果你把batch size提到16序列长度提到2048激活就会变成 16 × 2048 × 2048 × 154 ≈ 10.3GB瞬间成了一笔巨款。这也能解释为什么长序列训练那么吃显存——序列长度一涨激活内存几乎是线性甚至更快速地跟着涨。激活内存的经典缓解方案是gradient checkpointing也叫重计算。思路很简单forward时不全存中间结果只在某些节点存一个checkpointbackward到这个节点时再重新forward一次拿回中间结果。代价是多算一次前向换来的是激活显存大幅下降通常能砍掉一大半甚至更多。这在CS336后续实验里几乎是必开的选项因为你会真实感受到“同一个模型开和不开checkpointing能塞进去的batch完全不是一个量级”。3.3 算力与吞吐的快速估算显存只是一半资源核算还要算算力。同样大小的显存训练效率可能差很多因为还有算力峰值的利用率问题。Transformer训练里最常用的快速公式是单个step的计算量约为 6 × N × tokens。其中N是参数量tokens是这一个step里处理的有效token总数也就是 batch_size × seq_len。6的来源是前向每个token约2N次浮点运算反向约4N次合起来约6N。这个公式虽然省略了attention和LM head等细节但用于估算量级非常实用。举个例子1.3B模型配置batch_size32seq_len1024一次step处理约32768个token计算量约6 × 1.3e9 × 32768 ≈ 2.56e14 FLOPs。如果你的GPU单卡能达到约 1e14 FLOP/s已经是一个比较激进的有效算力那每步大概需要2.56秒。这个估算虽然粗但能让你在真正跑之前就判断一个实验大概要等多久是接单卡、几卡、还是等不了。4. 用PyTorch工具链实测显存从手动API到profiler账本算完了接下来需要工具来验证账本。PyTorch提供了从简单到复杂的各种测量方式。我按上手成本从低到高介绍一下。4.1 最简单的手动记账方式最直接的方法是调用torch.cuda下的内存统计函数memory_allocated()返回当前Tensor实际占用的显存已分配的活跃张量memory_reserved()返回当前进程向CUDA驱动申请到的显存总量。两者之差就是PyTorch显存分配器预留了但还没被使用的部分。测量峰值显存时先调用torch.cuda.reset_peak_memory_stats()跑完目标代码后用torch.cuda.max_memory_allocated()拿到峰值。我把这个逻辑封装成一个很简单的脚本站位下面的代码就是一个典型的参考模板import torch from transformers import AutoModelForCausalLM, AutoTokenizer model AutoModelForCausalLM.from_pretrained(your-model) model model.cuda().eval() torch.cuda.reset_peak_memory_stats() inputs torch.randint(0, 10000, (1, 2048)).cuda() with torch.no_grad(): logits model(inputs).logits current torch.cuda.memory_allocated() / 1024**3 peak torch.cuda.max_memory_allocated() / 1024**3 reserved torch.cuda.memory_reserved() / 1024**3 print(fcurrent{current:.2f}GB, peak{peak:.2f}GB, reserved{reserved:.2f}GB)注意这里用no_grad()包裹只是为了单独看推理路径。如果你想看训练路径就不要包no_grad并且要自己手动执行一次loss.backward()这样去看backward之后的峰值得到的数据才有参考价值。4.2 torch.profiler 看每个算子的内存开销手动API能看总量但定位不到具体是谁在占用显存。这个时候就要上torch.profiler了。它能按算子维度统计CPU/GPU时间、内存分配、显存使用等性能分析里最常用的一个组合是profile_memoryTrue。from torch.profiler import profile, ProfilerActivity with profile( activities[ProfilerActivity.CPU, ProfilerActivity.CUDA], profile_memoryTrue, record_shapesTrue ) as prof: # 这里放你想分析的训练step for step in range(1): loss model(input_tensor) loss.backward() print(prof.key_averages().table( sort_byself_cuda_memory_usage, row_limit20 ))输出表格里会列出每个算子自身的CUDA内存占用。按self_cuda_memory_usage排序后你能一眼看到最大的几块临时显存来自哪里。比如有时候会发现embedding层、attention dropout、或者某个softmax算子比想象中占得多。定位到具体算子之后你才有依据去替换实现、拆分模型、或者考虑更精细的并行方案。使用profiler时有个经验它本身会引入一些额外开销所以不要在长时间训练里一直开着最好只在个别step上采样。把训练跑慢一点没关系重点是拿到有代表性的内存数据。4.3 训练循环里如何自动记录资源曲线除了单次测量更实际的做法是在训练循环里按固定间隔记录资源消耗画出一条“显存曲线”。曲线比单个峰值更有诊断价值比如你可能观察到显存在前几个step稳定上升说明某些缓存或临时分配没有释放干净也可能看到optimizer.step()那一刻出现阶段性高峰说明优化器更新时临时状态特别大。我常用的记录方式是在每个training step的固定点打点def log_memory(tag): print( f[{tag}] allocated{torch.cuda.memory_allocated()/1024**3:.2f}GB, freserved{torch.cuda.memory_reserved()/1024**3:.2f}GB, fpeak{torch.cuda.max_memory_allocated()/1024**3:.2f}GB ) for step in range(total_steps): optimizer.zero_grad() loss model(batch) log_memory(after_forward) loss.backward() log_memory(after_backward) optimizer.step() log_memory(after_optimizer) torch.cuda.reset_peak_memory_stats()每隔几步reset一次峰值就能看到每个阶段正常情况下的资源波动。顺手验证一下你第3节手算的每步激活是否在合理范围内。如果算出来应该是1GB实际跑到5GB那多半有什么隐藏分配你没考虑到接下来就该拿profiler去查了。5. 藏起来的开销清单分布式、混合精度与运行时前面算的都是“账本上的常驻项”但实际训练里还有很多看起来不起眼、却能塞满显存的固定开销。我把最常见的几类列出来避免你算好之后被这些“隐藏项”打脸。先看一张常见隐藏开销速查表开销类型大致量级出现时机CUDA context数百MB到1GB进程初始化CUDA时cuDNN / cuBLAS workspace数十MB到数百MB某些GPU算子首次运行时NCCL通信缓冲每卡几十MB到数百MB初始化分布式通信时数据拷贝与padding临时张量取决于batch设计每次准备batch时混合精度的master weight4字节/参数DriveDDP或自有优化器保存时这里面最容易忽略的是CUDA context。很多人在单卡脚本里跑一个小模型一上来nvidia-smi就看到显存已经被占了几百MB甚至1GB第一反应是哪个程序泄漏了。其实这只是CUDA运行时初始化时预留的context属于固定成本。你无法取消它但可以在估算时预留出至少1GB的buffer。分布式训练里NCCL的通信缓冲是另一笔大开销。每个进程创建NCCL communicator时可能会为通信链路预留额外显存。卡数一多这笔buffer加起来也相当可观。再加上分布式数据并行在backward时会对梯度做all-reduce会把梯度打包成通信buffer这部分显存也是临时分配的很容易造成训练中期的突发峰值。所以做多卡训练时我习惯把预留buffer从40GB里先减个2-3GB再去做单卡需要的显存估算。另一个容易被忽略的是padding相关的临时张量。数据加载时为了把变长序列对齐到固定长度要生成attention mask和相关索引。如果padding逻辑写得比较粗糙可能会复制出好几份全尺寸张量。之前见过一个案例序列平均长度只有300但padding到1024之后batch里大量无效token把激活激活和attention矩阵全部撑大显存直接翻了快一倍。要想省这块开销最好的办法不是调显存而是在构造batch时就做动态长度排序和按桶分batch尽量让每个batch内部的长度接近减少padding比例。混合精度里还有一种容易被漏算的“账”虽然梯度往往以bf16存储但某些框架为了数值稳定性会保留fp32的梯度副本或者在loss scaling后做一次fp32的梯度转换。这些额外副本每个参数多占4字节别小看这4字节7B模型就是28GB。这就是为什么我不建议拿“16字节/参数”当成精确值的原因它只是让你快速判断量级的经验值真正的账要以你实际框架的存储设计为准。6. 一次真实OOM排查账本公式和profile工具怎么配合用最后分享一个实际案例正好能把前面的方法串起来。某次我在一张40GB的加速卡上跑一个1.3B模型的领域预训练实验。按第3节的估算bf16混合精度 Adam静态开销大约20GBmicro batch设为4seq_len 1024激活估算1.3GB左右。全部加起来不到22GB离40GB还有很大余量看起来完全没有问题。结果训练启动后还没走完第一步直接报OOM。当时的第一反应是“估算了这么低怎么还会爆”然后怀疑是激活公式用错了。但我没有直接去调小batch而是先做排查。第一步我把torch.cuda.reset_peak_memory_stats()放到训练循环最前面并且在报错前打印memory_allocated和memory_reserved。实际打印出来非常反直觉allocated只有25GB左右但reserved已经顶到40GB。这说明分配器从驱动那边申请了大量显存但大部分并不是“活跃使用的Tensor”。真正的麻烦是CUDA进程能申请的总显存上限已经撞到了墙哪怕还有空闲的reserved空间也可能因为现有分配被卡住而无法满足下一步的申请。第二步我用torch.profiler抓了一个step按self_cuda_memory_usage排序发现最大的内存消耗不是来自模型参数也不是常规激活而是某个attention算子在构建完整注意力矩阵时的临时输出。seq_len 1024看起来不夸张但attention矩阵的大小和批内序列长度、头数、头维度都有关系临时缓冲算下来远超我最初手动估计的激活公式。第三步我从两个方向改配置一是把micro batch从4降到2先把峰值压下来二是给模型开gradient checkpointing把attention里需要保存的大量中间激活重新计算而不是全部留在显存里。同时把序列长度从1024暂时降到768等到训练跑通后再逐步调回去。改动之后allocate峰值回落到26GB左右reserved维持在30GB以内训练能够顺利跑起来。事后复盘这次OOM的根本原因是我虽然算过静态账但没有把“框架运行时峰值”和“临时算子缓冲”纳入预算。账本公式帮我确认了方向但最终解决问题的是profiler定位到了具体算子。后来我养成了一个习惯第一次跑一个新的训练配置前固定先留出总显存15%-20%的buffer给运行时和通信再在关键节点用profile确认一次。这个习惯也让我意识到resource accounting不是一个一次性的动作而是一个持续校验的过程。刚开始训练时模型参数、梯度、优化器状态的账几乎不会变但激活部分和算子临时缓冲会随batch、序列长度、模型结构快速变化。每次调整配置之后与其靠感觉不如跑一两个step看看数据账本和测量结果对上了后面训练才安心。如果你也想把这套方法沉淀到代码里建议把显存记录封装成一个简单的context manager在训练脚本里随时调用每次完整训练结束后存一份资源日志。这样一来同一模型在不同配置下的显存行为就有了横向对比下次遇到OOM时你不是从一个空白状态开始排查而是直接调历史日志看哪个环节发生了变化。这个做法在我看来是整个resource accounting最值得长期坚持的部分。