最近在试着用 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 在7B这种规模上,第一次编译的显存峰值反而比正常跑要高不少,因为它要生成和优化图,得额外存中间表示和梯度相关的东西。我之前试过,如果不开mode,默认模式在编译阶段确实容易爆,但你可以试试先跑一个很小的profiling step,比如拿一个batch的1/4数据预热编译,等图优化完再上真实batch size,这样能避免峰值重叠。dynamic=True那个我理解是给动态shape用的,但如果你输入维度固定,开了它反而会触发多次重编译,时间当然更长,建议别加。至于省显存,我实际用下来感觉torch.compile对显存优化不是强项,它主要省的是kernel launch的开销,真正想降显存得配合gradient checkpointing或者把optimizer状态用bitsandbytes量化,这两个组合起来比单靠compile管用。另外你可以试下mode="max-autotune"配inductor的显存调度参数,比如设置triton的cuda cache上限,不过这个调起来比较玄学。
踩过同样的坑,24G卡编译7B确实容易直接OOM,我后来是先关掉dynamic,用reduce-overhead配合max-autotune的cudnn后端才稳住。另外你试试把编译范围缩小到单个decoder layer,别一下编整个模型,显存压力会小很多,速度提升也没差太多。graph break那个问题,我一般是把注意力里的reshape和view操作尽量统一,break少了编译开销才降得下来。对了,你微调用的是LoRA还是全参?全参的话建议别开compile,省那点时间不够折腾的。
说实话看到你这情况我第一反应是“太真实了”,torch.compile 对大模型微调来说不是无脑开的,它本质上是图编译的额外开销,尤其是第一次编译要跑 CUDA graph capture,那显存峰值会远超正常 forward,24G 被干爆很正常。我之前在 13B 模型上试过,只要开了 mode="default" 都会在编译阶段多占 6-8G 显存,后来学乖了,先用小 batch 把图编译出来,再调大 batch 跑,这样能绕开峰值爆炸。你说的 dynamic=True 让编译变长我也有同感,它本质上是把每个 shape 都当成动态去生成特化 kernel,数量多了自然慢,大模型微调里其实输入序列长度基本固定,你干脆别开 dynamic,把 padding 做死,编译一次缓存住后面就快了。关于 mode 的选择,我实际用下来 reduce-overhead 适合显存余量小的场景,因为它的 CUDA graph 优化幅度温和一些,max-autotune 是真的能提速但编译时间和显存占用都翻倍,你这种 7B 想省显存就别碰。还有个坑是 graph break,只要模型里有条件控制流或者数据依赖的 Python 操作,比如用到 dict 索引或者自定义的 RNN 类,就容易被 break,一旦 break 了编译几乎等于白做,建议先把模型里的动态 op 全换掉,比如把 attention mask 的构造放到模型外面。最后我自己的做法是,微调阶段干脆不编译,等推理阶段再开 torch.compile,因为训练时反向传播的梯度更新本来就会打断图优化,省的那点时间还不够补偿编译开销,你要是真想试,先把 batch size 调到 8 以下,再配合 max_memory 限制,应该能稳住。
同感,compile对大模型微调显存反而是负优化,试试关掉dynamic或换max-autotune能缓解点。
爆显存多半是编译时额外开销,先小模型调参数,别指望直接省显存。
遇到过同样的问题,7B模型用torch.compile默认模式直接给我把40G也干满了。后来发现关键是别开dynamic=True,那个会触发大量recompilation,显存峰值反而更高,编译时间翻倍很正常。我实际用下来,mode="reduce-overhead"配合max-autotune里的cudagraphs选项,再把batch size压到1,才能稳定跑起来,提速大概有15%但显存没省多少。另外你可以试试把编译范围限定在特定的子模块上,别整模型compile,这样graph break会少很多,显存占用也友好一些。你那个graph break具体报在哪几层?
试试max-autotune配dynamic=False,编译久点但显存峰值能压下来,不过7B还得靠gradient checkpointing兜底。
试试max-autotune加dynamic=False,编译久点但显存稳很多,reduce-overhead对7B真不够用。
max-autotune会疯狂试配置,显存肯定炸,先用reduce-overhead加dynamic=False稳住吧。
说实话你遇到的问题我上个月刚踩过一遍,24G跑7B微调本来就很极限,torch.compile的显存峰值反而比eager模式高不少,因为它编译期会生成额外的临时buffer,尤其第一次跑graph捕获那步,我这边直接多吃了3-4G。我个人建议是别指望它省显存,它的核心收益在推理或者大batch训练时的kernel fusion,微调场景你batch size卡死在1的话,提速也就10%左右,但显存风险反而更大。dynamic=True那个确实会让编译时间翻倍,因为它要生成多套shape的specialized kernel,对LLM这种序列长度固定场景基本没用,除非你的输入长度变化特别大。我现在的做法是直接用torch.compile的默认mode,然后配合activation checkpointing把显存峰值压下去,编译前先跑一个小的profile确认一下峰值到底爆在哪一步。另外graph break这个你帖子好像没打完,但如果是说编译时打断的warning,那个对显存影响不大,主要影响速度,可以试试把attention里的某些自定义op换成原生实现减少break。你要是真追求省显存,不如直接上gradient checkpointing加混合精度,compile放最后再考虑,毕竟它优化的是计算密度不是内存占用。
遇到过一样的坑,24G卡跑7B微调,torch.compile默认模式确实会先吃一波显存去做图捕获,跟你batch size没关系。建议先别上max-autotune,那个是拿显存换时间,试试mode="reduce-overhead"配合dynamic=False,编译时间能接受,显存峰值大概会降个10%左右。另外你提到的graph break,可以在编译前用torch._dynamo.config.suppress_errors=True看看有多少断点,断点多了编译收益反而低,不如直接不开compile用gradient checkpointing硬扛。对了,如果你用的是HuggingFace的trainer,记得关掉cache,不然编译缓存会重复触发,那也是显存杀手。
说实话,我试过几次之后直接放弃torch.compile了,大模型微调瓶颈主要在激活值显存,编译省的那点内存不够它编译过程折腾的。你要真想省,不如把注意力放到混合精度和gradient checkpointing上,batch size能翻倍。dynamic=True那个是给变长输入用的,你固定seq_len就别开,它每次shape变化都重新编译,时间全花在编译上了。graph break那边,LLM里attention和norm模块很容易触发,你要是看到一堆break提示,基本就可以放弃这个优化路径了。
我倒是觉得你可以先试一下把torch.compile放到推理阶段验证效果,微调场景真不一定划算
说实话我第一次跑7B也这样,compile的峰值显存比eager模式高不少,尤其max-autotune会疯狂试配置。我后来是先用reduce-overhead配合dynamic=False,把编译时间忍过去,等图稳定了再开dynamic,显存能压下来一些。另外你试试把gradient_checkpointing打开,跟compile叠加效果还行,别指望省太多,主要图个稳定。graph break那个问题,建议先看看是不是有动态shape或者自定义算子,我上次就是吃了这个亏,改成静态padding后编译快了一大截。
说真的,你踩的坑我上个月刚趟完一遍,24G卡跑7B微调,torch.compile默认配置基本就是给显存上刑。mode=“reduce-overhead”其实省不了多少显存,它主要是减少kernel launch的开销,适合小batch高吞吐的场景,大模型该爆还是爆。真正吃显存的是编译过程中的额外内存分配和guard检查,特别是dynamic=True会让编译器生成更多分支路径,显存碎片化更严重,编译时间翻倍很正常,我建议你直接别开dynamic,除非输入序列长度变化特别大。我现在实际在用的方案是配合gradient checkpointing,再加上max-autotune但是关掉cudagraphs,然后手动设置torch._dynamo.config.capture_scalar_outputs=False,这样能把编译时的峰值显存压下来不少。另外你注意下是不是把优化器的状态也一起编译进去了,那个很吃显存,可以试试只compile模型前向,反向和优化器留在原图。还有个偏方,先跑一次很小的warmup输入让编译器建好图,再换成真实batch,有时候能避开峰值分配。你提到的graph break其实很关键,微调时如果loss函数里有自定义操作,经常触发graph break,那torch.compile基本就退化成eager模式了,提速有限还白占显存,建议用torch._dynamo.explain看看打断点在哪。
我最近也在折腾这个,7B微调直接上torch.compile确实容易爆显存,尤其是默认模式它会为了图优化留不少额外缓存。你可以试试先不开dynamic,把reduction设为"max-autotune"但配合max-autotune-no-cudagraphs,cudagraph那部分经常是显存杀手。另外编译前把gradient checkpointing开了,batch size先压到1,等编译完再调大,实测能稳不少。graph break那个提示我也遇到过,多半是模型里有动态shape或者python控制流,先试着把input_ids的padding固定住,能少很多break。
torch.compile对大模型微调确实坑不少,我试过7B直接爆显存,后来发现关键是得配合reduce-overhead模式,而且batch size最好别动,用gradient checkpointing兜底。dynamic=True那个我也踩过,编译时间翻倍但收益微乎其微,感觉对小batch场景不如干脆关掉。另外你提到graph break,我建议先加torch._dynamo.config.suppress_errors=True跑一遍,看看哪些地方打断了图,手动改改那些动态shape的op,比硬调编译参数靠谱。你微调时是不是用了lora?如果是的话,可能得先冻结原参数再compile,不然显存会同时吃两份激活值。
说实话我踩过一模一样的坑,24G卡跑7B直接compile必爆,后来发现torch.compile本质是拿显存换速度,想省显存得先把max-autotune扔了,用默认模式然后配合gradient checkpointing才勉强压住。dynamic=True那个确实会拖慢编译,因为它要生成多个shape的specialized kernel,小模型无所谓,大模型纯纯折磨。我现在的做法是干脆不compile整个模型,只对attention或者MLP里的几个大算子单独compile,效果反而稳定。另外你提到的graph break,建议用torch._dynamo的日志看看到底断在哪,有时候是自定义算子或者python控制流导致的,修掉以后显存能降不少。
graph break那部分你大概率是遇到了,compile对动态shape特别敏感,dynamic=True会触发大量recompile,编译时间自然爆炸。我实际跑7B的经验是,先把max-autotune和reduce-overhead都试一遍,但真正省显存的核心是配合gradient checkpointing,compile本身优化的是kernel融合,不是内存分配。另外你的batch size=1的话,建议直接关掉dynamic,固定seq_len,让graph保持静态,不然光编译缓存就能吃满显存。还有个坑,别在forward里加Python条件判断,拆成两个模型分支反而更好。
我之前也被torch.compile坑过,24G卡编译7B属实有点极限,因为graph break和CUDA graph捕获会额外吃显存,尤其是dynamic=True会让编译器生成更多特化分支,缓存和临时tensor都变多,编译时间拉长是正常的。我自己的经验是,微调阶段别直接上compile,先用原生跑通一个小step,再用torch.compile(..., mode=“reduce-overhead”)试,那个模式其实对显存友好点,因为减少了kernel launch的开销,但不会像max-autotune那样疯狂做tiling搜索。还有个偏方,就是先compile一个小的子模块(比如只compile attention或者MLP),别整个模型一起上,显存峰值能降不少,速度也有提升。至于graph break,你输入里如果有动态shape或者Python控制流,它就会断成很多个小图,每个图都有一份中间缓存,这可能是你显存爆掉的主要原因。我后来干脆把dynamic设成False,固定sequence length,用padding凑齐,反而编译时间和运行显存都稳了。你提到max-autotune,那个真的别碰大模型,它会把所有候选kernel都试一遍,编译时间感人,显存峰值也高。建议你试试编译时加torch._dynamo.config.capture_scalar_outputs=True,有时候能减少一些意外的graph break,但别抱太大希望。最后想问下你用的是全参数微调还是LoRA?如果是LoRA,也许可以只编译base model的forward,不编译backward,能省不少。
踩过同样的坑,24G卡编译7B确实容易直接OOM。我的经验是先别开dynamic,默认模式编译一次后把优化后的模型存下来,之后加载用,能省不少编译时的临时显存。另外max-autotune确实比reduce-overhead更吃显存,但跑起来推理速度提升明显,微调场景其实用默认模式就够,主要省的是重编译开销。你试过把batch size再调小点,或者用gradient checkpointing配合吗?我这样配合下来,显存峰值能降30%左右。graph break那个问题,建议用torch._dynamo的日志看看具体在哪断的,有时候是某些自定义算子没兼容,改改代码结构反而比调参有效。
说实话你这个问题我太有共鸣了,之前调一个6.7B的模型也是被torch.compile搞得怀疑人生。第一次编译显存暴涨其实很常见,因为编译过程本身会生成额外的中间表示和算子变体,相当于多了一份临时开销,尤其是max-autotune模式会疯狂试配置,24G根本不够它折腾。我后来实践下来,大模型微调阶段其实不太推荐直接上torch.compile,因为forward和backward的图结构变化大,重编译频繁,省的那点时间全赔进去了。如果你非要用,我建议先用mode=“reduce-overhead”,这个模式对显存友好很多,而且别开dynamic=True,动态形状会让编译缓存失效,每次新shape都触发重新编译,时间自然爆炸。另外你可以试试把编译范围缩小,只compile transformer block里的attention部分,或者干脆冻结embedding和lm_head,这样编译图小很多,显存压力能降一个量级。还有个土办法,就是先用小batch size把编译跑完,等缓存生成后再切回大batch,虽然有点hack但实测有效。至于graph break,我遇到的情况是只要代码里有Python控制流或者自定义autograd.Function就容易触发,尽量把算子都改成torch原生操作,能减少很多麻烦。你现在用的PyTorch版本是2.1还是2.2?不同版本对编译内存的管理差别还挺大的。
试下mode="max-autotune"配合fullgraph=True,编译开销大但显存峰值反而低,dynamic=True如果输入形状固定就别开,动态shape会阻止很多fusion优化。另外开compile前先确保代码里没有Python-side的list拼接或dict动态操作,这些会强制graph break导致额外内存申请。我这边微调13B时是先把pad到固定长度再compile,batch_size开到2都没爆过。你那个graph break具体报在哪个算子?大概率是attention里的mask生成逻辑。
我之前也踩过这坑,torch.compile省显存的前提是能完整capture计算图,建议先把模型里所有条件分支和动态shape都改成静态的,比如把attention mask预先算好存起来。再者,7B微调可以试试把optimizer换