最近在试着用 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 条同款问题踩过坑,24G卡跑7B微调,torch.compile默认配置下CUDA memory直接涨了快4个G,后来发现是编译时给每个算子都预留了额外workspace,尤其是那种带reduce的融合kernel。我现在的做法是直接上mode="max-autotune"但配合memory_format=torch.channels_last,反而比reduce-overhead稳,显存占用和原生差不多,速度还快了20%左右。dynamic=True那个主要是给动态shape用的,微调时seq_len如果不变就别开,开了会触发大量重新编译和额外内存池分配,纯属给自己添堵。另外你提到graph break,这个才是显存炸的隐形元凶,建议先用torch._dynamo.config.log_code=True跑一次,看看是不是某些自定义层比如rope或attention mask写法导致打断图,打断后每个小片段都会单独做内存规划,碎片化特别严重。我后面把模型里所有in-place操作都改成函数式写法,graph break少了一大半,显存才稳下来。还有个小技巧,编译前先torch.cuda.set_per_process_memory_fraction设个0.9上限,至少能保进程不崩,方便你观察到底哪一步涨的显存。
torch.compile这玩意我一开始也踩了同样的坑,后来发现它默认的编译策略是“先完整构图再优化”,所以显存峰值反而比eager模式高不少,尤其是7B这种规模。你试试把mode设成“default”或者干脆用“reduce-overhead”,它其实会牺牲一点编译后的运行速度来换取更低的中间张量内存占用,batch size=1的情况下体感最明显。至于dynamic=True,那个主要是给输入shape频繁变化的场景用的,比如NLP里不定长序列,微调固定长度的话开了反而会让编译器放弃很多静态shape的优化机会,编译时间暴涨是必然的。另外有个小技巧是配合torch._dynamo的cache_size_limit调小一点,或者用torch.compile的fullgraph=False,把图拆碎,每个子图单独编译,内存峰值能降不少。不过说实话,大模型微调阶段我后来都直接关掉compile了,省那点时间不够折腾显存的,真要省显存还是得靠gradient checkpointing加混合精度。你那个graph break的问题,是不是模型里有动态控制流或者自定义autograd函数?那个确实会导致编译碎片化,得不偿失。
看到你说编译直接爆显存,我第一反应是太真实了。我拿13B模型试过,默认模式确实会额外吃掉不少内存,因为编译过程本身要保存图结构和中间张量,跟你微调时的激活值抢空间。我的经验是,先把max-autotune扔一边,那个调优搜索阶段跟个无底洞似的,动态shape开了就更夸张,等于每轮shape变化都要重新探索一遍,时间全耗在搜索上。
我自己跑通的做法是,先用mode="reduce-overhead"配合dynamic=False固定住输入长度,然后实在不行就分两步走——先用小batch把图编译缓存好,再切回正常batch跑,虽然有点投机取巧,但至少能避开第一次编译时的峰值。还有个小技巧,用torch._dynamo.config.suppress_errors=True先看看有没有算子在拖后腿,很多情况下是某个自定义op让编译退化成eager模式,反而更费显存。
不过我也有个没搞明白的地方,你提到graph break,是不是编译时报了具体哪些节点断掉了?我之前遇到类似情况是attention里用了太多python控制流,把那段逻辑改成纯tensor操作后,显存占用直接降了三分之一。你要不先把报错里的break点发出来,咱们对着具体位置调,比盲试参数靠谱得多。
试试max-autotune配合static_shape=True,小batch下能省不少显存,dynamic=True确实会增加编译开销。
max-autotune省显存纯属玄学,我试下来compile前先腾出10%显存做缓存才稳。你dynamic=True加长编译是正常的,小batch建议干脆别开。
graph break和显存占用是两码事,试试max-autotune配dynamic=False,能省不少内存。
试试max-autotune配合static shapes,dynamic=True对LLM反而拖慢编译还吃显存。
说实话reduce-overhead和max-autotune我都试过,对大模型来说峰值显存反而更高,因为编译过程会生成额外的中间张量,真正省显存得靠activation checkpointing配合,compile只是省了kernel launch的开销。dynamic=True那个主要是给动态shape用的,你固定sequence length的话开着确实纯浪费时间。我现在的做法是小batch先编译预热,跑通后再把batch加上去,不然一次编译直接把显存打满太常见了。graph break这问题也挺头疼,建议先看看有没有算子在拖后腿,把那些小op手动fuse一下可能比折腾mode参数更有效。
试试max-autotune然后配合gradient_checkpointing,dynamic别开,编译时间翻倍纯属浪费。
说实话你遇到的不是个例,7B模型直接compile对显存开销确实很猛,因为它默认会做很多保守的图优化和额外buffer分配。我自己的经验是,先别开dynamic,用默认模式加reduce-overhead,然后把batch size降到1再试,等编译完成后再慢慢调大,这样能避免一次炸掉。另外max-autotune千万别对大模型用,它搜配置的过程比训练还吃显存,纯属给自己找罪受。你提到的graph break问题,其实可以试试把模型里那些动态shape的部分(比如padding或mask)用torch.compile的fullgraph=True强制检查,虽然报错会多,但一旦跑通,省显存效果立竿见影。
说实话你这情况我太熟了,当时我拿torch.compile跑13B的模型也这样,第一次编译那会儿显存直接飙到比不开还高,后来排查才发现是CUDA graph的捕获过程和autotune的探针把激活内存暂时撑起来了。你提到的reduce-overhead其实主要省的是kernel启动开销,不是显存,它反而会为了图捕获多留一些缓存,所以大模型上我基本不用这个模式。max-autotune就更狠了,它会试一堆tiling策略,编译期间那显存占用简直像坐火箭,除非你batch size特别小并且能接受几十分钟的编译时间,否则真不建议在微调场景碰它。我自己现在用的办法是直接默认模式,然后配合torch._dynamo的配置把guard限制放宽一点,比如设置dynamic=False,因为对LLM来说shape基本固定,动态shape只会让编译器生成一堆分支代码,缓存和显存都更吃力。graph break那个问题也得多留意,如果模型里有自定义的Python控制流或者某些op不被支持,它会退回eager模式,这时候compile反而又慢又费显存,你可以用torch._dynamo.explain看看断点在哪,有时候把那些自定义层单独移到cpu上或者重写成torch原生op就能解决。另外一个小技巧是编译前先跑一个小的warmup batch,让显存分配器稳定下来,再上真实batch,这样峰值会稍微平缓一点。你要是微调的话,其实还可以考虑冻结底层几层只编译上层,或者干脆用gradient checkpointing搭配compile,虽然速度提升没那么明显,但至少不会一上来就爆。
这问题太真实了,我调7B的时候也踩过这个坑。torch.compile的显存开销主要在编译期的图捕获和额外缓存,建议先用mode=“default”跑通小batch,把dynamic关掉,你那个dynamic=True其实会触发多套shape的专门优化,编译自然慢到怀疑人生。我实际用下来,微调场景别指望省显存,最多是省点显存碎片,真正吃显存的大头还是activation和optimizer state,不如开gradient checkpointing实在。你说的graph break,如果能避免自定义算子就尽量别用,一break性能直接打五折。
24G跑7B微调本来就紧巴巴的,torch.compile的显存峰值反而更高太正常了,它编译期要存中间图和梯度,本质是拿空间换时间。你试试把dynamic关掉,固定shape后编译开销能小不少,还有batch size=1的话可以开reduce-overhead,减少kernel启动损耗,max-autotune更适合shape稳定且batch大的场景。另外可以看下是否触发了graph break,比如代码里如果有Python控制流或者动态shape,一break就退化成eager模式,显存和速度都白优化了。我实践下来大模型微调真要省显存,还是gradient checkpointing加混合精度更立竿见影,compile锦上添花,别指望它雪中送炭。
试试max-autotune配dynamic=False,编译久点但显存稳,或者干脆先关compile用gradient checkpointing顶一顶。
torch.compile这玩意对大模型确实不太友好,我试过7B微调,编译时显存峰值比正常forward还高,感觉它要存一堆graph中间状态。后来我是先关掉dynamic,用默认模式,然后配合gradient_checkpointing才压住显存,但提速也就10%左右。你那个24G爆掉,大概率是编译期的额外内存开销,不是运行时的,建议先小batch跑通编译再调大。另外max-autotune别轻易上,它搜一圈cuda kernel能把时间拉长好几倍,reduce-overhead相对稳一点。graph break确实烦,能尽量把模型里python控制流改成tensor操作会好很多,但LLM里有些自定义attention就没办法了。
说实话我第一次用torch.compile也这样,7B模型直接爆显存太正常了,这玩意儿编译期会额外占不少资源。你可以试试先不开dynamic,用默认模式把CUDA graph缓存关掉,或者干脆用mode="default"加reduce-overhead,省显存效果反而更明显。另外微调场景下torch.compile收益其实没那么大,尤其batch size小的时候,建议先检查是不是gradient checkpointing没开,那个比compile省得多。至于max-autotune,它会在编译时疯狂试配置,显存和耗时都爆炸,生产环境真不太推荐。