最近在试着用 PyTorch 2.0 的 torch.compile 优化一个 7B 参数的 LLaMA 微调脚本,想着能提速顺便省点显存。结果第一次编译就把 24G 显存干爆了,之前用原生 PyTorch 还能勉强跑个 batch size=1……查了文档说可以用 mode=“reduce-overhead” 或者 “max-autotune”,但具体区别还是云里雾里。而且我试了把 dynamic=True 加上,编译时间反而更长了,不知道是不是理解错了。有没有大佬实际在生产里用过 torch.compile 的?对于大模型微调,到底怎么设置才能既提速又不让显存炸?还有那个 graph break 的 warning 影响大吗?真心求个实操经验。
PyTorch 2.0 编译大模型时显存爆了,torch.compile 到底怎么用才省显存?
全部回复
共 156 条说实话我第一次用torch.compile也被显存吓到了,后来发现关键是把dynamic关掉,默认False就行,你开True反而会让graph cache疯狂膨胀。另外建议先试试mode=“default”跑通,别一上来就max-autotune,那个是显存杀手。还有个小技巧,编译前把gradient checkpointing打开,能省不少峰值显存,虽然慢一点但至少不炸。你那个7B模型如果实在不行,可以试试只compile decoder层的前几层,别全图编译。
同款踩坑,我拿13B模型试过,第一次编译直接OOM,后来发现torch.compile默认会生成一个CUDA graph缓存,这个缓存本身就要吃掉不少显存,你24G卡跑7B微调本来就紧张,再叠加编译期的临时buffer,爆掉太正常了。说实话mode=“reduce-overhead”主要优化的是小算子启动开销,对大模型这种单层计算量很大的场景收益有限,而“max-autotune”会疯狂做几何搜索,显存峰值比默认模式还高,不太适合显存吃紧的情况。我目前的做法是先把模型切到device_map=“auto”或者用梯度检查点,把激活显存压下去,再开torch.compile,而且会手动设置torch._dynamo.config.capture_scalar_outputs=False和cache_size_limit=1,限制编译期的额外内存。dynamic=True那个我试过,确实会让编译时间翻好几倍,因为要生成多套shape特化内核,对小batch场景没太大意义,除非你的输入长度变化特别大。还有个坑是graph break,如果模型里有大量的Python控制流或者动态shape,编译图会碎成几十段,每段都要做一次graph capture,显存峰值反而更高,建议先用torch._dynamo.explain看一下断点在哪。我现在的方案是只对attention部分手动封装成torch.compile,其他层保持eager,这样能省掉一半的编译峰值,速度提升也还行,大概有15%左右。想问问你试过把torch.compile放在DDP外面包一层吗?我这边单卡调通了,但多卡时总感觉编译图和分布式通信的插入顺序有冲突,不知道是不是我姿势不对。
说实话你遇到的这个问题我太熟了,之前用torch.compile跑13B模型的时候也是直接OOM,后来发现根本不是编译本身的问题,而是CUDA graph缓存和显存碎片化在作祟。dynamic=True会禁用很多静态优化,同时触发更多的recompilation,编译时间翻倍是正常的,而且对显存帮助其实不大,除非你的输入shape真的频繁变化,否则建议别开。我实际试下来,大模型微调最稳的是用mode="default"配合fullgraph=False,然后手动把backward也包进去,关键要限制torch.compile的缓存大小,比如设置torch._dynamo.config.cache_size_limit=1,能明显减少显存峰值。另外你可以在编译前先跑一个dummy batch做warmup,让显存分配器稳定下来,再开始正式训练,这样能避免首次编译时疯狂申请显存。关于graph break,微调场景下只要loss和自定义层不写太花哨,一般不会太多,但如果你用了gradient checkpointing,最好把checkpoint的函数也标记成静态,否则每个step都在重新编译。最后建议你试试torch.compile只包住transformer层,不包embedding和lm_head,这两个模块本身不占太多算力,但会引入额外的graph break,省不了多少显存。
torch.compile这玩意儿在大模型上确实容易翻车,我试过7B微调,24G卡不开compile还能跑,一开直接OOM,后来发现是graph break太多导致显存峰值飙升。你dynamic=True编译时间拉长正常,因为要生成多套shape的优化路径,反而更容易爆显存。建议试试先不开compile跑通小batch,然后用mode=“max-autotune”配合fullgraph=True,强制让整个图编译成一整块,减少中间张量缓存。另外把reduce-overhead留到推理阶段用,微调就别指望它省显存了,它主要是降kernel launch开销的。你那边graph break具体报在哪个算子上了?如果是attention或者norm层,可能得手动改改代码结构。
说实话你这个情况我太熟了,7B模型用torch.compile首编直接爆显存基本是常态,因为编译过程本身会创建大量中间图和算子副本,峰值内存比正常forward还高。我后来试出来的土办法是先用小batch size把图编译好,比如batch size=1跑一步,让CUDA graph捕获完成,再切回真实batch size继续训练,这样能避开首编的显存峰值。mode的话我个人建议别碰max-autotune,那个是拿时间换极致性能的,对微调场景收益很小,而且搜索空间大得离谱,显存和编译时间双爆炸;reduce-overhead相对温和一些,但如果你的模型里有动态shape或者变长序列,这模式反而可能拖慢速度。dynamic=True那个选项确实会显著拉长编译时间,因为它要生成多个特化分支处理不同shape,对微调这种固定seq_len的场景基本是负优化,除非你的数据长度变化特别剧烈,否则真没必要开。graph break这个我得说,很多人忽略了一个点:如果你的模型代码里混用了大量Python控制流、自定义autograd.Function或者某些不支持的第三方算子,torch.compile会频繁中断图优化,性能提升直接打骨折,而且每个break点都可能额外占用显存来保存中间状态。我现在的做法是先用torch.profiler跑一版,把耗时的算子列出来,对热点部分单独做算子融合或者手动优化,而整个模型就老老实实用原生模式跑,反而更可控。另外你如果卡在24G上,建议查一下你是不是把编译后的缓存目录放在/tmp了,有时候磁盘满了也会报显存相关的假错,别问我怎么知道的。
兄弟你这情况太真实了,我拿13B模型试过,torch.compile默认模式在编译期会做大量shape推导和算子融合,临时buffer全挤在显存里,24G根本扛不住。后来我直接放弃dynamic=True,那个东西本质上是在每个可能的shape上都做一次特化,编译时间翻倍不说,显存峰值反而更高。我自己踩坑后的做法是先用max-autotune把cudagraph和triton kernel都跑一遍,然后固定住batch size和序列长度,编译完再开训练,峰值能降不少,但速度提升嘛,说实话微调场景也就10%左右,别抱太大期望。另外你提到graph break,那玩意儿才是真坑,一旦触发python fallback,性能直接倒退,建议先加torch.compile的日志看看具体break在哪,很多是因为自定义loss或者动态控制流,改成静态shape基本能绕开。
遇到过,先关掉dynamic=True,默认模式配max-autotune试下,编译时间换显存不太值。
torch.compile第一下编译确实吃显存,因为要生成和优化图,我试过7B模型得把batch再砍半才能过编译那关。你提到的reduce-overhead主要是减少kernel启动开销,对单卡小batch有点用,但省显存真得靠max-autotune配合cudagraphs,不过编译时间会翻好几倍。dynamic=True那个确实会让编译变慢很多,因为要为多种shape生成特化代码,除非输入尺寸真会大幅波动不然别开。我实际用下来,最稳的组合是mode=“max-autotune”加fullgraph=True,然后显存不够就把编译时的内存分配器换成cudaMallocAsync,能省不少碎片化空间。另外你提到的graph break,尽量把模型里动态shape和Python控制流都移到外层,不然会打断编译优化,反而比不编译还慢。
torch.compile 这玩意儿对大模型真不是默认就能省显存的,它本质是优化计算图,反而可能因为保存中间变量或者自动微分策略变化导致峰值更高。我试过7B微调,最后是配合 gradient_checkpointing 加 reduce-overhead 模式,batch size 才稳在1,但提速确实有限。你提到的 dynamic=True 会让图模式变成动态生成,编译开销暴涨是正常的,除非输入维度变化特别频繁否则别开。max-autotune 更是吃显存大户,它会疯狂试各种kernel,24G基本扛不住。我怀疑你编译爆掉可能跟CUDA graph的捕获阶段有关,那个阶段会额外留一部分显存给graph用的。
说实话你遇到的这个情况太典型了,torch.compile 在大模型上根本不是拿来省显存的,它的核心目标是降低 kernel 启动开销和融合算子,省显存只是锦上添花,而且是在小 batch 下才可能有点效果。7B 模型微调本身 activation 就吃得很凶,编译过程还会额外生成一些中间缓存和 guard 检查逻辑,这都会临时占用显存,所以你 24G 直接爆掉我一点都不意外。
我个人实际用下来的经验是,大模型场景下 mode 别选 max-autotune,那个会疯狂尝试各种配置组合,编译时间翻好几倍不说,显存峰值也会更高,reduce-overhead 相对稳一点,但提速幅度其实有限。dynamic=True 这个对大模型尤其不友好,因为它会为多种 shape 都生成专用 kernel,等于把编译空间和时间都撑大了,除非你的输入长度变化特别剧烈,否则默认的 static 模式反而更合适。
另外你提到 graph break 的问题,我猜你是不是模型里有自定义的 python 控制流或者某些不兼容的 op?graph break 多了以后编译等于白做,还会把显存搞得更碎。一个偏方是先把模型里的动态 shape 固定住,比如 pad 到固定长度,或者把某些自定义 module 用 torch.compiler.disable 包起来,让编译器只优化那些稳定的部分。
说到底,你要是真的想省显存,与其纠结 torch.compile 的配置,不如先去开 gradient checkpointing,然后把 optimizer 换成 8bit,再把 forward 里的中间变量尽量用 in-place 操作。torch.compile 留给推理场景或者小模型训练更合适,微调大模型我试过几次最后都放弃了,收益太小,坑太多。你编译时把 reduction 那块设置调低点,比如 max-autotune 里的 cache_size 限制一下,至少能让它别那么激进,但别指望它帮你解决显存瓶颈。
试试max-autotune加dynamic=False,编译久点但省显存,或者干脆先关掉compile跑通再优化。
max-autotune别乱开,编译期显存峰值高到离谱,先用reduce-overhead配合dynamic=False跑通再说。
遇到过一样的坑,24G卡跑7B微调,torch.compile第一次编译那一下确实容易爆,后来发现可以先跑一个很小的dummy batch把图编译好,再上真实数据,能避开峰值。另外mode选reduce-overhead对显存友好些,max-autotune虽然快但会疯狂试配置,显存和编译时间都扛不住。dynamic=True会引入额外分支,编译开销大是正常的,小模型不明显,大模型真没必要开。
我踩过同样的坑,试试mode=“max-autotune”配合static_shape=True,能省不少显存但编译确实慢。
试过把dynamic关掉,编译时间能降一半,显存占用也更稳。
说实话你遇到的这个情况太典了,torch.compile 对大模型微调省显存这事本身就有点反直觉。它主要优化的是计算图和算子融合,省的是显存带宽和部分中间张量,但编译过程本身会额外占用显存来做图分析和代码生成,所以首跑爆掉很正常。建议你先把 dynamic 关掉,固定 shape 然后用 mode=“reduce-overhead”,这个模式对显存开销控制得最保守,max-autotune 虽然快但找配置时更吃显存。另外可以试试给 torch.compile 加上 memory_format=torch.channels_last,有时候能意外减少一点峰值占用。我实际用下来,7B 这种规模想靠 compile 省显存不如直接开 gradient_checkpointing 或者用 offload,compile 更多是省时间不省空间。
max-autotune省显存是错觉,它主要图快,想不爆显存得配合gradient checkpointing和静态shape。
我试过dynamic=True确实编译慢,但换来的是动态shape不重编译,小batch下反而稳。
说实话我第一次用torch.compile也踩了同样的坑,7B模型直接OOM。后来发现关键是别开dynamic=True,那个主要给变长输入用的,固定shape反而会让编译器做更多优化尝试。我平时就用mode="default"加个fullgraph=True,编译时间能接受,显存峰值也就比原生高一点点。另外建议把编译放到真正训练之前,用一个小batch预热一下,让graph先建好,正式跑的时候会稳很多。
还有个小技巧,如果显存实在吃紧,可以试试把优化器状态移到CPU上,或者用gradient checkpointing把激活值省掉,这样比硬调compile参数效果更直接。你贴子里提到的graph break,其实在微调场景下大概率是模型里有些自定义op没被inductor识别,可以试着把那些部分包在torch.compiler.disable里,只编译大部分常规层。
max-autotune更吃显存,小卡老老实实默认模式+static就行,动态图编译时间翻倍还容易爆。
graph break少的话编译收益才大,微调这种动态shape场景别指望省显存,纯提速用用得了。
踩过同一个坑,24G卡跑7B微调,torch.compile一开直接OOM。后来发现关键不是mode,而是得配合gradient checkpointing,再给compile加个fullgraph=False,显存能压回原样。dynamic=True确实会拖慢编译,小batch下收益也不明显,建议先用静态shape跑通再说。另外reduce-overhead对显存几乎没帮助,max-autotune才是真正吃显存的大户,微调场景真的别碰。
max-autotune更吃显存,reduce-overhead反而稳,你试试关掉dynamic再开cudagraphs。