最近在试着用 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 条我试过mode=“reduce-overhead”确实能省点显存,但batch size还是得调小,编译那步本身就吃资源。
同感,我试7B模型也是compile直接爆显存,后来发现加mode="reduce-overhead"配合dynamic=True反而更吃内存,可能是图缓存的问题。建议先关掉dynamic,用默认模式只对前向做编译,或者试试把batch size降到1再开编译,我这样勉强能跑起来。另外graph break确实是个坑,得尽量保证计算图完整,不然分段编译反而更慢。
我也是被torch.compile坑过好几回,特别是大模型这块,真的不是无脑套上就能省显存的。你提到的mode="reduce-overhead"其实更适合小模型或者推理场景,它主要减少Python开销,但对大模型那种计算密集型的任务,显存优化并不明显,甚至可能因为编译临时缓存多占一点。max-autotune就更狠了,它会穷举各种融合策略,编译时间和显存峰值都会飙升,7B模型用这个基本就是自爆。我自己的经验是,微调场景下直接开torch.compile默认模式,配合torch.inference_mode()或者梯度检查点(gradient checkpointing)一起用,反而更稳。另外dynamic=True确实会让编译时间变长,因为它要处理动态输入形状,如果你的数据长度基本固定,建议别开,否则每次变形状都可能触发重新编译。还有一个容易忽略的点:编译前先把模型用half或者bfloat16加载,能有效降低初始显存占用,避免编译过程中间缓冲溢出。至于graph break,通常是模型里有些操作不被torch.compile支持,比如某些自定义CUDA核或控制流,可以试着用torch.compiler.disable装饰一下出问题的子模块,或者干脆只编译Attention部分。说实话,现在torch.compile对大模型微调的生产落地还不太成熟,我自己更倾向先用deepspeed或者zero3做显存优化,等模型跑稳了再考虑用compile做锦上添花。
同感,第一次用torch.compile我也被显存吓了一跳。后来换成mode=“reduce-overhead”确实好点,但7B模型想省显存还是得配合gradient checkpointing和更小的batch size。dynamic=True主要是为了适应可变输入长度,如果你的数据形状固定其实没必要开,开了反而增加编译开销。另外可以试试把编译范围缩小到模型的前向部分,别一股脑全编译,能缓解不少压力。
这问题太真实了,torch.compile 对显存的开销经常被低估,它默认会生成一些辅助buffer和额外算子,峰值内存反而比eager模式高。你试试把dynamic=True去掉,配合max-autotune + mode="reduce-overhead",但别对这个组合抱太大希望,微调场景下编译收益往往不如推理明显。另外建议开gradient_checkpointing,并且把编译范围缩小到模型的核心decoder层,别整个模型一股脑塞进去。你那个graph break如果指的是打断次数,可以试试用torch._dynamo.config.suppress_errors看看具体是哪里被fallback了,很多时候是自定义的loss或attention实现导致的。
torch.compile省显存得配着max-autotune和静态shape用,dynamic=True反而让显存峰值更高。
说实话我第一次用torch.compile也这样,7B模型直接OOM太正常了。后来我干脆只在推理阶段开compile,训练时候还是老老实实用原生模式,省心很多。你要是非要在微调时用,建议把dynamic关掉,那个主要是给输入shape会变的场景准备的,LLM里基本用不上,开着反而增加编译开销。另外可以试试把编译范围缩小到某个子模块,比如只compile attention部分,不要整模型一把梭。graph break那玩意儿确实烦,能避开就避开,实在不行就看看是不是有什么动态控制流在捣乱。
建议直接关掉dynamic,用reduce-overhead配合gradient checkpointing,省显存比编译实在多了。
这题我熟,之前调13B模型时也踩过这坑。torch.compile本身不会省显存,它主要是算子融合和减少kernel launch开销,省显存得靠reduce-overhead配合gradient checkpointing,或者干脆把batch再砍半。dynamic=True确实会拖慢编译,因为要生成动态shape的fallback路径,小模型无所谓,大模型纯属自找麻烦。另外你提到的graph break,如果代码里有太多python控制流或者自定义autograd,建议先换成torch原生算子,不然编译图碎成渣,比不编译还慢。
踩坑经验+1,torch.compile在7B这个量级上确实容易先吃一波显存来换编译期的优化,尤其是max-autotune会疯狂试配置。我后来是先用mode=“default”把图跑通,确认显存峰值没问题后再考虑要不要上reduce-overhead,dynamic=True那个是给动态shape用的,你固定batch size的话开了反而会生成多套kernel,编译时间和显存都更糟。另外如果你用的是LoRA微调,其实可以只compile主干部分,把adapters排除掉,效果会好很多。graph break那个问题我也遇到过,建议先查一下是不是有自定义op或者Python控制流卡住了trace,实在不行就配合torch._dynamo的日志看看哪里断了。
24G炸了太正常了,我拿13B试的时候直接OOM到怀疑人生。你先把dynamic=True关了吧,那玩意本质是给输入shape频繁变化的场景用的,LLM微调序列长度基本固定,开了只会让编译器疯狂做guard检查,编译时间翻倍不说,显存峰值反而更高。我实际跑下来最稳的组合是mode="max-autotune"加fullgraph=True,但前提是得把模型里那些动态控制流全清掉,比如Python的if或者list append,不然graph break一多,优化等于白做。另外你提到graph,我猜你想说graph break?这个确实是大坑,建议用torch._dynamo.explain看一下具体断在哪,大部分情况都是因为用了tensor.shape[0]这种动态维度去控制循环,改成固定batch或者用torch.where把条件逻辑向量化能解决不少。至于省显存,torch.compile本身不是用来压显存的,它主要省的是kernel launch的开销,显存峰值反而可能因为cudagraph缓存而涨一点,想省显存还是得靠gradient checkpointing加混合精度,compile跟这俩配合好了才有效果。我现在的做法是先把模型freeze住,只开需要训练的那几层,再用max-autotune编译,24G跑7B lora微调勉强能塞下,速度大概提升15%到20%,但编译那几分钟是真肉疼,不过一次编译后缓存住,后续迭代就快了。你试过把torch.cuda.set_per_process_memory_fraction设到0.9没?有时候是碎片化问题,限一下反而能跑。
说实话你遇到的这个情况太典型了,torch.compile 在 7B 模型上第一次编译会额外消耗大量显存,主要是因为它要生成和优化计算图,中间会留很多临时缓冲区。我自己的经验是,别指望它帮你省显存,它的核心收益是省时间,显存开销通常反而会涨10%到30%,所以你要么把 batch size 再调小,要么干脆用 gradient checkpointing 配合着来。
关于 mode 的选择,我劝你别上 max-autotune,那个会疯狂试各种 kernel 组合,编译时间翻几倍不说,显存峰值也会更高。reduce-overhead 相对温和,但说实话在微调场景下收益也不明显,因为瓶颈往往在数据传输和反向传播上。dynamic=True 这个我也踩过坑,它会让编译器假设输入形状多变,导致生成更多分支代码,自然又慢又吃显存,如果你的数据 shape 固定,千万别开。
另外你提到 graph break,这个才是关键——如果你的模型里有大量 Python 控制流或者自定义 autograd.Function,torch.compile 会频繁打断图优化,性能甚至可能退化。我建议你先把模型里那些不常用的分支简化掉,或者用 torch._dynamo.config.suppress_errors=True 看看到底哪里 break 了。
最后,对于大模型微调,我实际用下来最稳的组合是:CUDA graphs 手动开起来(比 torch.compile 省心),配合 static input shape,再加上 activation offload 到 CPU。你试试把 compile 只包在 transformer 层上,别包整个模型,这样编译时间和显存峰值都会降不少。你用的是 HuggingFace 的 LLaMA 实现还是原生代码?如果是 HF 的,记得把 gradient_checkpointing 和 torch.compile 一起用,不然反向传播时显存会直接翻倍。
踩过一样的坑,compile 默认会做 graph break 和算子融合,但编译期本身就要额外显存来存中间表示,7B 模型直接爆很正常。我后来是先用 reduce-overhead 模式跑通,再把 dynamic 关掉,编译时间能接受,显存峰值大概只比原生高一点点。max-autotune 是真别碰,它会疯狂试配置,24G 根本不够造。另外你可以试试把模型切到 bf16 再 compile,或者先 freeze 掉不更新的层,能省不少临时 buffer。graph break 那个问题,我建议开 torch._dynamo 的日志看看具体断在哪,有时候是某个自定义 op 不兼容,换个写法就好了。
说实话我第一次用torch.compile也这样,7B模型直接爆显存很正常,因为编译过程本身会额外占用内存做graph捕获和优化。建议先试试mode="default"或者"reduce-overhead",别一上来就max-autotune,那个会疯狂试配置。dynamic=True确实会拖慢编译,因为它要处理动态shape,你如果固定batch size和seq len就别开。另外可以试试把编译范围缩小,只compile attention或者某些层,别整个模型都包进去,这样显存压力会小很多。我自己用下来感觉torch.compile对大模型微调的收益主要在推理阶段,训练时除非batch特别大否则提速有限,你可以先分阶段测试下。
说实话你这个情况我太熟了,当时我拿torch.compile跑13B的模型做lora微调,也是24G卡直接OOM,后来干脆把compile范围缩小到只包住那些计算密集的attention和mlp层,其他部分原样跑,显存瞬间就稳住了。你提到的mode参数,reduce-overhead主要就是减少kernel launch的开销,适合小算子多的网络,但大模型里真正吃显存的是中间激活值,这个mode帮不上什么忙;max-autotune会疯狂试各种tiling和融合策略,编译时间暴涨是正常的,而且有时候找出来的最优配置反而更吃显存。dynamic=True那个我理解是给动态shape留缓冲,但代价就是编译器不敢做太多静态假设,生成代码更保守,显存占用可能不降反升。我现在的做法是直接不开dynamic,固定seq_len和batch,然后配合gradient_checkpointing一起用,compile编译出来的graph会被checkpoint切碎,反而能压住峰值。还有个偏方是先把模型用torch.compile编译完再load权重,别在load之后compile,顺序错了显存会多出一份临时副本。最后想问你一下,你编译时那个graph break有没有报warning?我这边老是有一些自定义op触发break,一break整个编译就退化成eager了,提速效果直接归零,如果你也遇到这个,可以试试把那些op改成纯torch原生算子,虽然麻烦点但至少能保证graph完整。
你这情况太真实了,torch.compile 在编译期本来就会额外吃一波显存,尤其 7B 模型直接上 max-autotune 基本是自杀。我建议先别开 dynamic,固定 shape 用 reduce-overhead 试试,编译内存会小很多,跑起来也稳。另外 graph break 我怀疑是你代码里有些动态控制流或者自定义 op 没处理好,可以先看看编译日志里 breakdown 有没有 warning,把那些地方改成静态写法。省显存真正有效的是配合 gradient checkpointing,compile 只是提速,别指望它当显存救星。
同款踩坑,24G卡跑7B微调,compile一开直接OOM,我后来查了issue发现这玩意默认会做完整图捕获,内存峰值比eager模式高不少,尤其动态shape的时候。你试试把dynamic关掉,然后mode用“default”或者干脆不指定,让torch自己选保守策略,我这么改之后显存基本回到原生水平,速度反正微调场景提升有限,主要是图优化对固定shape的推理收益大。另外graph break那个问题,建议看下编译日志里有没有break的提示,如果频繁break那基本等于白编译,还会多占内存。我现在是干脆只在推理阶段用compile,训练阶段老老实实eager+gradient checkpointing,省心得多。你要是非要训练时用,可以试试把编译范围限定在某个子模块上,别整个model一起compile,这样内存可控很多。
踩过同样的坑,24G卡跑7B微调本来就紧巴巴的,torch.compile的显存开销主要来自图捕获和额外缓存,建议先别上max-autotune,默认模式配合static_shape=True试试,你的batch固定为1就不该开dynamic,那会触发大量重编译反而更慢。另外可以给编译函数包一层torch.no_grad(),或者把输入padding到固定长度,能明显减少显存峰值。至于graph break,它其实不影响显存,但会拖慢速度,可以试着把模型中频繁变shape的层单独抽出来不编译,其他部分再走compile。
这问题我太有同感了,7B模型直接compile确实容易把显存干穿,我当时是把max-autotune换成默认模式,再配合gradient_checkpointing才勉强跑起来。dynamic=True那个主要是给变长输入用的,你固定batch size的话开了反而增加编译开销,建议关掉。另外graph break不一定全是坏事,有时候break了反而能省点显存,你可以看看编译日志里到底断在哪,说不定是某些算子不支持导致的。
试过max-autotune反而更吃显存,建议先用reduce-overhead配合static_shape试试,dynamic=True确实会拖慢编译。
图模式可以试试关掉cudagraphs,有时候省显存比提速更实际。