最近在试着用 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 到底怎么用才省显存?
全部回复
共 10 条
碰到过类似的问题,7B模型用torch.compile默认模式确实容易显存暴涨,我试下来mode=“reduce-overhead”对显存友好很多,但提速有限。dynamic=True主要是给变长输入用的,如果序列长度固定的话开着反而增加编译开销,建议关掉。另外可以试试先对单层做compile测试,别一上来就整模型,graph break的问题我还没完全搞明白,蹲个懂哥讲讲。
说实话,你这个batch size=1都爆显存的情况我太熟了。torch.compile在LLM上的显存开销主要来自两个地方:一个是编译过程中生成的中间缓存和重计算图,另一个是CUDA graph的预留显存。你直接调默认设置,它会把整个forward+backward都做一次全图捕获,7B模型光这个图结构就能吃掉好几个G。
关于mode的选择,我的经验是“reduce-overhead”其实更适合推理场景或者小batch训练,它默认会帮你做CUDA graph,但graph本身就会锁住一大块显存作为workspace。做微调的话,我建议你试试mode=“max-autowrap”或者干脆用mode=“default”然后配合torch._dynamo.config.capture_scalar_outputs=True,这样能减少一些不必要的张量特化。dynamic=True会让编译器保留多个图的版本,编译时间当然会翻倍,但对微调来说dynamic其实挺关键的——如果序列长度或者batch size有变化,没有dynamic的话每次变化都会触发重新编译,那个时间损失更大。
还有一个很多人踩的坑:torch.compile默认会对所有算子做重计算优化,这个在大模型上反而容易让显存爆炸。你可以试着在torch.compile外面套一层torch.cuda.empty_cache(),或者用torch._inductor.config.fx_graph_cache来缓存已编译的图。另外,如果实在显存吃紧,可以考虑用torch.compile只包裹attention部分而不是整个模型,这样编译收益还在,但显存压力小很多。你那个graph br后面是不是想说graph break?那个确实是动态图的痛点,建议你先用torch._dynamo.explain看看哪里断了。
我也碰到过一模一样的情况,24G显存直接炸穿,后来发现torch.compile默认模式其实对显存并不友好,尤其大模型第一次编译会做图捕获和算子融合,临时占用的显存比跑推理还高。我目前试下来,先用mode=“default”配合dynamic=False,至少能让编译过程不爆显存,虽然速度提升有限但至少能跑起来。你提到的reduce-overhead理论上会减少框架调度开销,但实际对大模型微调来说,显存节省并不明显,反而编译时间更长。max-autotune就更夸张了,它会穷举所有算子配置,显存峰值能飙到40G以上,完全不适合24G卡。我自己的做法是先用unset_torch_compile_cuda_graph禁掉CUDA graph缓存,然后手动设置torch._dynamo.config.cache_size_limit到一个很小的值,这样编译过程中的中间缓存会被及时清理。另外dynamic=True确实会让编译变慢,因为它要为不同shape生成多个优化路径,但如果你微调时sequence length经常变化,这个开关反而能避免反复重编译,算是用时间换灵活性。你最后提到的graph break问题,我猜是遇到了算子不支持编译的回退情况,可以试试加torch.compile(fullgraph=True)来强制报错,定位到底是哪段代码没被捕获,然后手动替换成兼容写法。
同感,我也在折腾torch.compile,我的7B模型直接爆了32G显存,后来发现dynamic=True其实是为了动态shape场景设计的,如果输入shape固定就别开,不然编译开销会翻倍。另外我试过mode="max-autotune"确实更吃显存,不如先试试默认模式或者reduce-overhead,听说有人配合gradient checkpointing能压到20G以内。你那个graph break报的具体是什么算子?是不是自定义的attention模块没被捕获?
我也在纠结这个问题,试过用mode="reduce-overhead"跑6.7B的模型,编译完第一次前向确实省了点显存,但后面反向传播又炸回去了。dynamic=True那个我之前也踩过坑,感觉它主要是给输入shape变化大的场景用的,固定batch size反而拖慢编译。有没有试过给torch.compile传个backend="aot_eager"或者限制一下graph的切分策略?我也蹲个大佬讲讲生产上怎么调参。
graph break的问题确实让人头疼,我试过把model改成eval模式或者禁用某些自定义的autograd Function能减少一部分break,但效果不稳定。显存爆炸这个我猜主要是torch.compile在编译阶段会额外分配一些临时buffer来做图优化,尤其是max-autotune模式会跑很多次尝试找最佳配置,大模型上24G肯定扛不住。建议先关掉dynamic=True,它会让编译器保留更多动态路径,反而增加内存开销,reduce-overhead模式相对轻量,可以先试试这个。另外如果batch size已经小到1还爆,可以看看是不是模型里有什么动态shape或者控制流在torch.compile下被放大了。
老实说我也踩过这个坑,torch.compile 对大模型优化其实挺看硬件和场景的。mode="reduce-overhead" 更适合小模型或推理,对大模型编译反而容易因为图太大导致显存爆炸,而 "max-autotune" 就更吃显存了。建议你试试先关掉 dynamic=True,它对动态 shape 友好但编译开销巨大,微调时固定序列长度反而更省显存。另外 graph break 可能是因为代码里有 Python 控制流或自定义算子,试着把模型 forward 里能静态化的部分单独抽出来,或者用 torch._dynamo.config 里的 suppress_errors 先绕过不兼容的节点,实测能缓解不少。
试下 mode="reduce-overhead" 配合 torch.inference_mode(),亲测能省个几G显存,编译时间也能接受。
说实话碰到同样的问题,24G卡编译7B模型确实容易炸,我试下来mode=“reduce-overhead”相对稳一点,但提速有限。dynamic=True会增加重编译次数,显存占用反而变高,小batch size下不建议开。建议先关掉编译跑一次看峰值显存,再手动调torch._dynamo.config.cache_size_limit和张量分片策略,或者试试用fsdp+compile组合,能压一压显存。对了,graph break主要看有没有算子不支持动态shape,可以加torch.compile的日志排查下具体是哪块断了。
试下把dynamic关掉,开reduce-overhead配合梯度检查点,能压不少显存。