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_memory=True会让DataLoader使用锁页内存来临时存放数据,这样从CPU复制到GPU的速度会明显变快。但注意,锁页内存属于系统内存而不是显存,它不会增加你的显存占用,却会显著增加宿主机的内存消耗。如果你的服务器本身内存不多,把num_workers拉高,再加上pin_memory=True,可能模型还没开始训练,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_size=2048,num_layers=24,seq_len=1024,micro 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时不全存中间结果,只在某些节点存一个checkpoint,backward到这个节点时再重新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_size=32,seq_len=1024,一次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(f"current={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_memory=True。
from torch.profiler import profile, ProfilerActivity with profile( activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA], profile_memory=True, record_shapes=True ) as prof: # 这里放你想分析的训练step for step in range(1): loss = model(input_tensor) loss.backward() print(prof.key_averages().table( sort_by="self_cuda_memory_usage", row_limit=20 ))输出表格里会列出每个算子自身的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, " f"reserved={torch.cuda.memory_reserved()/1024**3:.2f}GB, " f"peak={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 weight | 4字节/参数 | 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,静态开销大约20GB;micro batch设为4,seq_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最值得长期坚持的部分。