最近在试着用 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 条max-autotune别硬上,7B微调先用reduce-overhead,dynamic除非输入变化大不然别开。
试过max-autotune确实吃显存,但配dynamic=False能压住,编译时间也短不少。
试过max-autotune反而更吃显存,建议先关掉dynamic,小batch多步验证再上全量。
碰到同样的问题,24G卡跑7B微调,torch.compile一开直接OOM,后来发现其实跟mode关系不大,主要是编译期会额外分配一堆graph相关的缓存。你可以试试先关掉dynamic=True,那个确实会拖慢编译,而且对微调场景收益很小。另外建议把torch.compile包在推理或特定层上,别整个模型都编译,或者用reduce-overhead配合更大一点的batch试下,有时候省的不是显存而是内存带宽。还有那个graph break,大概率是模型里有动态shape或者自定义op,可以用torch._dynamo.config.log_level查一下具体打断点在哪。
遇到过同样的问题,24G卡跑7B微调,torch.compile第一次编译那下确实像在渡劫。我后来是先用原生模式把数据load和preprocess跑通,再单独对forward部分做compile,别一上来就整model.compile。dynamic=True那个主要是给动态shape用的,你固定batch size的话没必要开,开了反而增加编译开销和显存峰值。另外你可以试试先跑一次小batch的warmup,把编译缓存生成好,再上真实batch,能躲过那个峰值。还有graph break这东西,能不触发就不触发,有时候为了省显存手写一些控制流,结果反而让编译碎片化,更吃显存。
graph break那块我建议先别纠结,7B模型直接上max-autotune就是找爆显存,先用默认mode把编译跑通再说。另外dynamic=True确实会拖慢编译,因为要生成多个shape的优化版本,微调场景下输入长度变化大的话不如直接padding到固定长度。我之前跑13B时是配合gradient checkpointing加reduce-overhead,batch size开1,显存刚好卡住没炸,速度提升大概有15%左右。你试试把编译范围缩小到某个子模块,别全模型compile,效果会好很多。
试试max-autotune加dynamic=False,编译前先小batch预热,显存能压不少。
遇到过类似情况,后来发现把gradient checkpointing打开,compile才真正省显存。
这问题我太懂了,7B模型用torch.compile首编就是会吃满显存,因为graph capture要额外存一份CUDA graph上下文。你试试编译前先用小batch warmup一下,或者干脆只对decoder layer做compile,别整模型一起上。dynamic=True确实会拖慢编译,因为它要预留各种shape的kernel,非动态shape场景真没必要开。max-autotune我实测提速有限还费显存,reduce-overhead在微调场景反而稳一点。另外你提到的graph break,建议把attention里那些自定义算子尽量换成原生实现,break一多省显存效果直接打骨折。
踩过一样的坑,24G卡跑7B微调本来就很极限,torch.compile的显存峰值反而比eager模式高,因为编译过程要额外存图。我后来是先用mode=“default”把编译跑通,再切reduce-overhead,另外把dynamic=False固定住,动态shape会让编译器疯狂recompile,显存直接翻倍。你那个graph break的问题,可以试试在模型forward里加torch._dynamo.mark_dynamic,或者干脆把embedding和lm_head这些层排除在compile之外,能省不少。目前我用下来,编译主要省的是kernel launch的时间,显存真没省多少,想省显存还是得靠gradient checkpointing。
max-autotune编译时确实吃显存,建议先用reduce-overhead跑通,另外dynamic=True别乱开,小batch下纯属负优化。
max-autotune会疯狂试配置,显存肯定炸,先试试reduce-overhead加dynamic=False再说。
踩过同样的坑,24G卡跑7B微调,compile一开直接OOM太真实了。我的经验是先用max-autotune配合fullgraph=True,但别开dynamic,那玩意对静态shape的微调场景纯属负优化。另外建议把编译范围缩小到具体的decoder layer,别整个模型一起compile,这样显存峰值能降不少。还有个野路子是先用小batch试编译,成功后再把batch加回去,有时候是编译时的临时buffer在作祟。至于graph break,我到现在也没完全搞懂,但发现把attention里的某些自定义kernel换成官方实现能减少不少,你可以试试。
试试把dynamic关掉,编译时先留足临时缓存,或者干脆用reduce-overhead模式,省显存还得靠梯度检查点。
遇到过同样的问题,7B模型上torch.compile的显存峰值反而比eager模式高不少,尤其是第一次编译时CUDA context和graph缓存会吃额外显存。我的做法是先用手动gradient checkpointing把激活显存压下来,再开compile,mode用reduce-overhead就好,max-autotune虽然理论上更快但编译时间和显存都太激进。dynamic=True确实会拖慢编译,因为要生成更多特化分支,小batch下不值得。你后面被截断的graph break问题,可以试试把模型里动态shape的部分(比如attention mask)固定下来,或者用torch._dynamo.config的suppress_error选项先跑通再说。
补充一点,如果显存实在紧,可以先关掉compile做完整训练,只在推理阶段开torch.compile,这样能避免微调时的显存峰值,速度提升也很明显。另外7B模型建议开torch.cuda.amp混合精度,配合compile比单开compile效果更好。你那个graph break具体报的什么错?有时候是自定义layernorm或者rope导致的,换回原生算子就解决了。
碰到同样问题的人不少,你这不是个例。torch.compile 本身不是拿来省显存的,它主要优化计算图和算子融合,反而在编译期会额外占用显存和内存,尤其 max-aututune 会疯狂试配置,7B 模型没做梯度检查点或者 offload 的话,24G 直接炸太正常了。我实际试下来,大模型微调想用 torch.compile,得先把显存余量留够,比如 batch size 再降一半,或者开 activation checkpointing,不然编译那一下的峰值根本扛不住。
关于 mode 的选择,reduce-overhead 其实更适合小 batch 或者推理场景,它减的是 kernel launch 的开销,对训练帮助有限。max-autotune 在训练里收益也不大,反而编译时间长到怀疑人生,除非你跑很多 step 能把编译成本摊薄,否则微调这种几百步的任务根本不划算。dynamic=True 那个我建议你直接关掉,它会让 graph 变成动态 shape 的,每次输入尺寸变化都可能重新触发重新编译,时间当然长,而且对大模型来说 shape 基本固定,开了纯属给自己找麻烦。
我自己现在微调 7B 级别的模型基本不用 torch.compile,除非是 A100 80G 这种显存特别宽裕的卡,而且会用 torch._dynamo 的 cache limit 控制一下编译重试次数。如果真想提速,不如先试试 bf16 混合精度加 flash attention,这两样对显存和速度的改善比 compile 直观多了。另外你提到的 graph break,我猜是 graph break 导致的回退到 eager 模式,那样编译就白做了,可以用 torch._dynamo.explain 看下 break 原因,多半是某个算子不支持动态 shape 或者自定义 op。你是在用 peft 的 LoRA 吗?如果是的话,可能有些层被 freeze 了导致 graph 构建出问题,可以试试只编译 trainable 的部分。
这题我熟,去年调13B模型的时候也被compile坑过。你试试把dynamic关掉,微调时shape基本固定,开dynamic反而让编译器疯狂做特化,显存和编译时间都爆炸。另外max-autotune对7B这种规模性价比很低,我实际用它比reduce-overhead多吃了快3G显存,速度提升却不到5%。建议先用reduce-overhead+static模式跑通,显存实在不够就配合gradient_checkpointing,compile对checkpointed部分的优化其实挺友好的。还有个坑是编译前先把input shape固定好,别用None维度,不然graph break会特别碎。
说实话你这个情况我太熟了,7B模型上torch.compile第一次跑编译期那一下就是会给你来个“惊喜”,24G直接被打满很正常,因为编译的时候要生成多个中间表示还得做图优化,临时显存峰值比你正常前向要高不少。我自己的经验是别一上来就开compile,先跑通一个step确认数据流没问题,再用torch.compile(model, mode="default")去试,max-autotune那个模式虽然能榨出极限性能,但编译时间跟显存开销都太猛了,微调场景下性价比很低。另外dynamic=True这个flag确实会让编译变慢,因为你要给各种动态shape都生成优化分支,如果训练时seq_len和batch_size固定,就别开了,反而拖后腿。至于显存爆的问题,我建议配合gradient_checkpointing一起用,把激活重算开起来,再把optimizer换成分片的,比如bitsandbytes的8bit adam,这样能腾出不少空间给编译的临时buffer。还有个坑是graph break,如果你模型里有大量python控制流或者自定义算子,编译时图会断成好几段,反而增加显存碎片,这种情况不如直接用torch._dynamo.disable装饰掉那几个模块,剩下的部分再compile。你现在是光编译就炸,还是跑到训练step中间炸的?如果是前者,试试把compile放到模型加载之后、数据进卡之前,同时设一下torch.cuda.set_per_process_memory_fraction(0.8)限制一下峰值,至少能先跑起来看看真实收益。
说实话你这个问题我踩过一模一样的坑,7B模型用torch.compile直接爆显存太正常了,因为编译过程本身会创建额外的中间图和优化缓冲区,峰值内存比实际推理要高不少。我后来试下来,对大模型微调最稳的是先用mode=“default”或者干脆不开编译跑通流程,确认显存够用再逐步加编译选项,别一上来就追求max-autotune。关于dynamic=True,那个主要是给输入shape频繁变化的场景用的,微调时序列长度固定的话真没必要开,开了反而会触发多次recompile,编译时间翻倍不说,显存峰值也会更高。你要真想省显存,不如把注意力放到gradient checkpointing和混合精度上,这俩对显存的改善比compile直观得多。另外我有个疑问,你编译爆显存的时候是只编译了forward还是把整个训练step都包进去了?如果包了整个step,试试只对推理部分做compile,训练部分保持原样,有时候能避开最吃显存的那个编译阶段。graph break那个问题,建议去看一下编译日志里具体break在哪个算子,一般就是自定义loss或者动态控制流导致的,把这些地方用torch.compiler.disable包起来反而效率更高。
遇到过一样的坑,24G卡跑7B直接OOM太真实了。后来我把dynamic关掉,固定shape,编译时间短了显存也稳了点,动态shape对graph break影响很大。另外可以试试把optimizer和model一起compile,或者干脆只compile forward部分,反向保持原样,能省不少峰值显存。max-autotune别轻易开,它找tiling配置那步本身就吃显存,reduce-overhead其实够用了。你可以用torch.profiler看下具体哪块爆的,有时候是编译缓存没清干净。
24G显存跑7B微调本来就很极限了,torch.compile的graph break会触发回退到eager模式,反而增加内存峰值。我试过把max-autotune和dynamic=True一起开,编译时间能绕地球三圈,但显存没省多少,后来干脆用reduce-overhead加static shape,batch size压到1才勉强稳住。你这情况不如先试试gradient checkpointing,把activation存下来换显存,编译带来的收益其实没那么大。另外可以查一下是不是CUDA graph缓存没清干净,有时候显存碎片也会导致爆掉。