最近在试着用 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那个reduce-overhead模式主要是减少kernel启动开销,但编译时显存峰值高挺正常的,因为它要生成和优化图。我试过7B模型,建议先把max-autotune放一放,这个模式会疯狂试配置,显存和编译时间都爆炸。dynamic=True确实会慢,因为要保留动态shape的fallback路径,小batch下不如直接静态。你试试配合gradient_checkpointing,然后编译前先warmup一个小输入,把cudagraph缓存起来,能缓解不少。graph break那个问题,可以先看看是不是有Python控制流或者自定义op卡住了,拆成小模块单独compile试试。
踩过同样的坑,compile第一次编译的CUDA graph缓存和显存峰值确实吓人,我后来是先跑一个小batch的warmup把图建好,再切回正常batch size,能缓解不少。dynamic=True那个主要是给变长输入用的,微调里如果序列长度固定,开着反而白增加编译开销,关掉试试。另外max-autotune会疯狂试配置,显存和编译时间双重爆炸,生产里我一般用默认mode加个inductor的显存上限设置,省心很多。顺便问下你用的是哪个版本的PyTorch?2.1之后对LLM的编译支持改善挺明显的。
踩过同一个坑,编译期显存峰值高是正常的,因为graph break和CUDA context初始化会额外吃显存。我后来是先关掉dynamic=True,用默认模式跑通一次生成CUDA graph,再开reduce-overhead做推理,微调时反而别开这个。另外max-autotune会疯狂试配置,显存肯定爆,建议用小模型先探路。你试试把torch.compile放在dataloader之后、模型forward之前调用,别包整个训练循环,可能能缓解一点。
说实话我第一次用torch.compile也这样,7B模型直接爆显存太正常了,因为编译过程本身会额外开一堆内存来做图优化。我后来是先用max-autotune跑小batch把cache存下来,再切回reduce-overhead跑正式训练,显存能压下来不少。dynamic=True确实会拖慢编译,除非你的输入shape频繁变化,不然别开,它会让编译器放弃很多静态假设。另外你提到graph break,我建议去torch._dynamo里看看有没有warning,能避免的break尽量用torch.compile的fullgraph参数强制检查一下,有些自定义算子改改写法就能绕过去。
说实话第一次编译爆显存太正常了,我试过7B的模型,torch.compile默认会生成多个guard分支,graph break又会导致额外内存开销,建议先把dynamic关掉,用mode="reduce-overhead"配合static=True,batch size保持1跑通再说。max-autotune那个是给推理场景用的,训练时tune时间太长还不一定省显存,反而可能因为cudagraph缓存占用更多。另外你可以试试把编译范围缩小到单个block上,别一上来就整model.compile,或者干脆用torch._dynamo的skip参数把某些ops排除掉,这样编译时间短很多,显存压力也小。
我这边实际测试下来,reduce-overhead模式配合max-autotune或者inductor的自动调优,显存反而比不编译高一点,但速度提升明显,关键是要把环境变量TORCHINDUCTOR_FORCE_DISABLE_CACHES设上,不然每次编译都缓存成新图,越跑越吃显存。你那个dynamic=True更吃显存是因为它要维护动态shape的特殊化分支,等于多存了好几份图,建议直接关掉。另外graph break的问题,可以试试把模型里那些自定义的python控制流用torch.compile的guard_fn参数绕过,至少能减少一半编译开销。
我倒是觉得,7B这种规模用torch.compile纯属给自己找麻烦,
跟你一样踩过坑,24G卡跑7B用compile基本是自杀式操作。我后来是先用原生模式把batch调到能跑,再开reduce-overhead,显存大概多占10%但速度确实上来了。dynamic=True那个主要是给动态shape用的,你微调输入长度固定的话别开,编译时间翻倍还容易爆。还有graph break的问题,建议代码里尽量少用python控制流,把能合并的op放一起,不然分段编译更吃显存。
torch.compile对7B模型微调确实容易在编译阶段吃满显存,因为它要生成和优化计算图,临时缓冲区开销很大。我之前试过把dynamic关掉(默认False)反而更稳,dynamic=True虽然能应对动态shape但会触发更多重编译,时间暴涨很正常。建议先用mode="default"配合max-autotune-no-cudagraphs跑通,cudagraphs在微调场景下经常因为显存不足直接炸。另外可以试试把编译范围缩小到单个Transformer层,别整个模型一起compile,或者干脆先用gradient_checkpointing把激活显存压下来再编译。你那个graph break是在哪个算子报的?如果是自定义attention,大概率是flash attention或者rope那块没被捕获,得手动改写成torch原生算子才行。
说实话你这个情况我踩过一模一样的坑,7B模型用torch.compile默认配置确实会先吃一波显存做graph capture,尤其是编译期间临时张量特别多。我后来是先用max-autotune配个memory_format=channels_last,再把dynamic关掉,batch size固定住,编译时间虽然长点但跑起来显存反而比原生省了大概15%,提速也明显。你试reduce-overhead如果还爆,可以看下是不是CUDA graph缓存区占的,设个torch._dynamo.config.cache_size_limit=1能压不少。另外graph break如果出现在attention里,可以试试把模型拆成几个子模块分别compile,别整体包,这样编译压力小很多。
遇到过一模一样的坑,24G卡跑7B微调,torch.compile一开直接OOM,当时差点以为是我代码写崩了。后来排查下来发现,编译本身会额外申请一块graph内存池,加上CUDA context和cuDNN benchmark的缓存,峰值显存比原生模式高出一大截,这跟你batch size多大没关系,纯粹是编译期的临时开销。你现在如果只是想要不爆显存,先把mode设成“default”,别碰“max-autotune”,那个会疯狂试kernel组合,显存和编译时间都爆炸。“reduce-overhead”主要优化的是小kernel的启动延迟,对7B这种大tensor计算其实收益有限,而且它也会增加显存占用。
dynamic=True那个我试过,确实会显著拉长编译时间,因为它要生成多个shape的specialized kernel,除非你的输入序列长度变化特别频繁,否则对大模型微调来说得不偿失。我现在的做法是先用torch.compile的“default”模式跑通,然后配合gradient checkpointing把激活显存压下来,等训练稳定了再开“max-autotune”做profile,看哪个kernel真的值得优化。另外你提到graph break,这个很关键——如果模型里有动态控制流或者Python原生list操作,compile会频繁打断图,导致生成的代码反而不如eager快,你可以用torch._dynamo.explain查一下打断点在哪,把那些部分用torch.compiler.disable包起来,让它们走eager,其他部分保持编译。
还有一个偏门技巧,编译前先把torch.cuda.set_per_process_memory_fraction设到0.9,强制让编译期显存分配有上限,虽然可能触发重试,但至少不会直接OOM崩掉整个进程。最后建议你盯一下峰值显存,用torch.profiler看内存分配事件,很多时候是某个临时tensor在编译优化时被复制了一份,手动fuse掉那些操作能省不少。
遇到过同样的问题,24G卡跑7B微调,torch.compile一开直接OOM,后来发现根源是编译期会额外申请显存做图优化,跟实际推理阶段完全两码事。我现在的做法是先用原生模式把batch和gradient checkpointing调到能跑的状态,再用mode="default"配合dynamic=False,只在数据形状固定时开编译,反而稳很多。至于reduce-overhead,它主要减少kernel启动开销,但省显存效果真不明显,max-autotune更是吃显存大户,大模型慎用。另外你提到dynamic=True变慢是正常的,因为要生成多个特化版本,微调场景下动态shape收益不大,建议直接关掉。还有个坑是graph break,建议把模型里容易触发break的算子(比如自定义attention)先重写成torch原生实现,否则编译退化不说,显存还多占。
说实话reduce-overhead和max-autotune在7B这个量级上收益真没那么大,编译本身吃掉的显存和内存反而更扎眼。我之前试过把dynamic关掉,固定序列长度,用max-autotune配合gradient checkpointing,batch size=1勉强能压到20G以内,但编译时间直接翻倍。你这情况不如先别上compile,把注意力放到flash attention和混合精度上,显存省得更多还省心。graph break那块我没太看明白你卡在哪,能说具体点吗?
踩过一样的坑,compile对大模型微调真不是默认参数就能直接上的。我试下来max-autotune省显存效果最明显,但编译时间能等到怀疑人生,reduce-overhead更稳一点。dynamic=True建议别开,除非你序列长度变化特别大,不然纯纯增加编译开销。另外记得把memory_format设成channels_last,配合compile能省不少显存。graph break这问题我也遇到过,核心是尽量避免在forward里用data-dependent的control flow,把动态shape的逻辑挪到外面去。
torch.compile省显存是玄学,我这7B直接关掉dynamic反而稳,max-autotune编译久但跑起来真香。
max-autotune别硬上,先试试reduce-overhead配合dynamic=False,编译时间能接受再调。
试试max-autotune加dynamic=False,编译慢点但显存稳,7B用24G别开reduce-overhead。
同款踩坑,24G卡跑7B微调,compile前勉强能跑,compile一开直接OOM。后来发现主要是编译期会额外保留一些中间张量做图优化,建议先试试mode="default"关掉reduce-overhead,那个模式为了省kernel启动开销反而会加大显存峰值。dynamic=True别开,大模型输入shape本来就固定,开了等于逼编译器生成多套特化代码,内存翻倍还拖慢编译。另外可以看下是不是gradient checkpointing和compile冲突了,我这边把checkpointing放到compile外面才正常,显存基本持平但速度确实提了20%左右。
说实话reduce-overhead那档基本就是给推理场景用的,训练时省显存效果很有限,我试过7B模型开这个反而多了几个G的峰值。max-autotune确实能压一些内存,但编译时间长得能去泡杯咖啡,而且必须配合静态shape用,你开dynamic=True等于白折腾。建议先别急着上compile,把gradient checkpointing和8bit优化器开了,batch size调到能跑的最大值,再考虑用compile做锦上添花。另外你提到的graph break,其实大部分情况是模型里有动态控制流,比如llama的kv cache更新,试试把那段逻辑改成静态shape能不能缓解。
我前段时间也在7B上踩过这个坑,一开始以为torch.compile是银弹,结果编译期直接给我把显存吃到OOM,后来发现关键其实不在mode,而是得配合gradient checkpointing一起用,不然graph本身就要占掉不少activation内存。mode=“reduce-overhead”主要是减少kernel launch的开销,对显存帮助有限,max-autotune会疯狂试配置,编译期长到怀疑人生,而且显存峰值反而可能更高,因为要存更多中间buffer。dynamic=True那个我也试过,它会让编译器生成更通用的kernel,所以编译时间暴涨,但对小batch来说收益真不大,除非你的输入形状一直在变,否则别开。我现在是这么干的:先开gradient checkpointing,然后用mode=“default”或者干脆不开compile,只把forward里的某些大算子手动融合,效果反而更稳。另外你可以试试torch._dynamo的cache_size_limit调小一点,有时候编译缓存太多也会把显存吃满。graph break那个问题我也遇到过,建议用torch.compile的fullgraph=True强制不切图,虽然编译失败会报错,但至少能告诉你哪里有问题,比黑盒强多了。
另一个思路是干脆用FSDP或者DeepSpeed ZeRO-3把模型参数分片,这样torch.compile只编译单个rank上的子图,显存压力会小很多,但通信开销得自己权衡。我最后是放弃了在微调阶段用torch.compile,只在推理阶段用,毕竟训练时反向传播的graph太复杂,收益真没想象中高。
graph break 才是关键,先试试把 dynamic 关了,用 reduce-overhead 加 max-autotune 组合,显存不够就上梯度检查点。
说实话reduce-overhead在7B上真不一定能省显存,它主要减少的是Python开销,编译期反而会多存不少中间张量。我之前试过,batch size=1的话干脆别开dynamic,那个会触发多次重编译,显存峰值更高。你不如试试把编译范围缩小到backbone的某几个大block,或者干脆用fsdp+compile配合,单卡24G想跑7B微调本来就很极限。另外注意下gradient checkpointing要和compile一起开,有时候能压下来几个G,但速度提升可能就没那么明显了。