最近在试着用 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 条试过把batch size再调小或者关掉gradient checkpointing试试?我这边用reduce-overhead配合static shape显存反而稳一些。
老实讲,你遇到的这个问题我当初也踩过坑,torch.compile 在大模型上真不是无脑套用就行的。你提到的 mode=“reduce-overhead” 其实主要是减少 kernel launch 的开销,对显存优化帮助不大,而 “max-autotune” 会疯狂尝试各种融合策略,编译时间和显存爆炸都是家常便饭。我个人在 7B 模型微调时试下来,比较稳的做法是先用 torch.compile 配合 dynamic=False,然后设置 mode=“default” 或者干脆不设 mode,虽然提速没那么猛但至少不会直接 OOM。另外 dynamic=True 那是给输入形状频繁变化的场景用的,你微调如果固定 batch size 其实没必要开,开了反而让编译器多生成一堆备选图,编译时间长而且显存占用更高。还有你提到的 graph break 问题,可以试试在编译前把模型里那些动态控制流(比如 if-else 或者循环)尽量用 torch.jit.script 或者直接写成静态张量操作,避免 graph 被切碎导致额外内存开销。最后一个小建议——如果你显存实在吃紧,可以考虑只对模型的前几层或者注意力部分做 compile,后面留给原生 PyTorch,这样既能吃到部分加速又不会让显存直接炸掉。
我最近也踩过这个坑,torch.compile对大模型确实有点暴力,默认模式会尝试全图优化直接把显存撑爆。建议你试试mode="reduce-overhead"加上dynamic=False,编译时间能接受,显存占用比原生还低一点。另外可以配合gradient checkpointing一起用,我这边的7B模型batch size从1提到了2,速度也有提升。那个dynamic=True主要是给动态输入形状用的,固定batch size的话没必要开它。
试下mode=“reduce-overhead”加dynamic=False,我这么调7B模型显存从24G降到20G了。
试试把memory_format设成channels_last,batch size再调小点,dynamic=True对静态图反而拖慢速度。
老实说torch.compile在7B模型上翻车太正常了,我个人经验是先用mode=“default”跑通小batch验证编译正确性,再慢慢调batch size,别一上来就全量编译。dynamic=True确实会拉长编译时间,因为要生成多个计算图版本,对小模型还行,大模型上反而得不偿失。另外graph break是个隐藏坑,建议用torch.compile的日志输出查一下哪些算子没被捕获,手动优化那部分代码比硬撑着整个图编译要稳得多。
实际试过7B模型,torch.compile确实是个显存大户,尤其是第一次编译时CUDA图会吃很多临时显存。我自己的做法是先设mode="reduce-overhead"关掉一些激进优化,再配合gradient checkpointing把激活显存压下来,这样batch size=2勉强能跑。dynamic=True那个主要是给动态shape用的,固定batch size的话开它确实会拖慢编译,反而得不偿失。另外graph break可能跟代码里某些Python控制流有关,建议先确保模型前向里没有if/for那种动态路径。
老实说我也被torch.compile坑过,7B模型用默认模式直接显存溢出太真实了。我后来试了mode=“reduce-overhead”确实能省点显存,但编译时间还是长,而且dynamic=True对动态shape有用,但静态batch下反而增加开销。建议你先关掉dynamic,试试gradient checkpointing配合,或者把model partitions拆得更细一点,graph break多的时候手动优化一下算子融合。
说实话你这个问题太真实了,我一开始用torch.compile也被显存搞懵过。mode="reduce-overhead"主要减少的是Python解释器的开销,对显存占用影响其实不大,真正吃显存的是编译过程中产生的大量中间图和临时缓存;而"max-autotune"会做更多算子调优,编译时间和显存消耗都会飙升,对小模型还行,7B参数量真心不建议。dynamic=True这个参数很多人理解有偏差,它不是直接省显存的,而是让编译器能处理动态输入形状,代价是编译时生成更多图分支,所以编译时间变长是正常的。我自己的经验是,对大模型微调,可以先试试默认模式加个torch._dynamo.config.cache_size_limit=64之类的限制一下缓存,或者手动把batch size调得更小、用梯度累积来撑过编译阶段。另外你提到的graph break问题很关键,如果模型里有大量动态控制流或者自定义算子,编译时会被打断成多个子图,不仅省不了显存反而可能变慢,建议优先检查forward里有没有if/for语句或者Python内置函数。说到底,torch.compile对大模型的收益目前还不稳定,我见过不少case是提速了但显存涨了20%-30%,不如先开着torch.inference_mode或者用原生AMP加gradient checkpointing来得稳妥。
试下mode=“reduce-overhead”配合dynamic=False,我这边7B模型显存占用降了30%但编译时间确实长了点。
老实说,你遇到的这个问题太典型了,我一开始用torch.compile的时候也是被显存直接干懵了。其实它默认的编译模式是会为每个可能的shape都生成一份优化后的图,对于大模型这种参数规模,编译过程中间变量和梯度缓存叠加起来非常恐怖。我试下来,如果你只是微调不是训练从头开始,建议先把mode设成“reduce-overhead”,同时一定要配合max-autotune的搜索限制,不然它真的会在编译阶段疯狂试探显存上限。还有dynamic=True那个参数,它本质上是让编译器为动态shape准备多个优化路径,所以编译时间暴涨是正常的,如果你能固定batch size和序列长度,尽量别开这个。另外有个小技巧,你可以先用torch.compile装饰部分子模块而不是整个模型,比如只编译attention层,这样显存压力会小很多,而且推理阶段提速效果依然明显。不过我也挺好奇,你在用gradient checkpointing吗?那个配合torch.compile有时候会冲突,得调一下环境变量才能稳定。
说到torch.compile的显存问题,我这边也踩过差不多的坑。其实核心在于编译本身会生成额外的中间图和缓存,对大模型来说这些临时结构反而可能把显存撑爆。你提到的mode参数我试下来,reduce-overhead确实比默认模式更保守一些,但省显存效果有限,max-autotune就更耗了,它是在找最优kernel而不是省资源。dynamic=True会让编译器为不同shape都做优化,所以编译时间暴涨是正常的,除非你输入尺寸变化特别大,否则一般不建议开。我实际的做法是先用torch.compile的backend=“aot_eager”试试,它不做图优化但省编译开销,或者干脆把编译范围缩小到模型中的某些大模块,比如只compile attention部分。另外有个偏方是编译前先跑一次dummy input让显存分配稳定下来,再调小batch size慢慢试。你那个graph break的问题,可以检查下代码里有没有动态控制流或者自定义操作,尤其是LLaMA的attention mask和position embedding容易触发断点,导致编译器退化成eager模式反而更吃显存。
说到这个我可太有感触了,7B模型用torch.compile,默认模式下编译过程本身就会吃掉大量临时显存,尤其是第一次编译做graph tracing的时候,相当于把整个计算图摊开做优化,24G确实扛不住。我个人的经验是,微调场景下别指望compile能省显存,它主要是为了加速重复计算,反而在编译阶段会额外占显存。你可以试试先关掉compile,用原生torch跑一个step来预热显存分配,或者直接用mode=“reduce-overhead”并配合torch._dynamo.config.capture_dynamic=False,把动态shape的追踪关掉,这样编译负担会小很多。dynamic=True那个选项其实更适合输入长度变化很大的推理场景,微调里输入长度相对固定,开了反而让编译器多做无用功。另外,如果你不是必须用图优化,干脆对模型的后半部分单独compile,比如只对attention层做编译,其他层保持eager模式,这样能平衡速度和显存。你提到的graph break我猜是编译过程中因为某些操作不支持而分段,分段越多编译开销越大,可以检查下模型里有没有自定义算子或者控制流,尽量换成torch原生算子来减少断点。
跑7B用reduce-overhead模式加batch size=1,再把dynamic关掉,显存能稳不少,编译慢是正常的。
老实说我也踩过这个坑,24G显存跑7B模型用torch.compile确实容易直接炸掉,尤其是默认模式会把整个计算图都做一次完整编译,对显存开销特别不友好。我后来试下来,mode用"reduce-overhead"比"max-autotune"稳很多,后者会疯狂试各种fusion策略,编译时间直接翻倍,显存峰值也更高。dynamic=True那个参数我建议小模型或者输入长度变化不大的场景再开,大模型开它反而会让编译器做更多guard检查,编译时间暴涨,而且对于微调这种固定batch size的任务收益很小。还有个trick是配合torch._dynamo.config.cache_size_limit调小一点,或者干脆用torch.compile的fullgraph=False,只编译部分子图,牺牲一点加速比但能保住显存。我目前生产里是先用profile跑一轮,找出最耗时的子图单独编译,剩下部分用eager模式,这样batch size能回到原来的一半左右,速度也有10-15%的提升。你试过把torch.compile作用在单个transformer block上而不是整个模型吗?这样显存压力会小很多,而且调试起来也直观。
老实说,你遇到的这个情况太常见了,torch.compile 虽然宣传得很香,但实际对大模型微调确实有点“暴力”——它默认会把整个计算图做一次全量编译,7B 模型光那个图展开就够喝一壶的。我试过 mode="reduce-overhead" 和 "max-autotune" 的区别,前者其实是减少 kernel launch 的开销,对显存友好一些,但编译时间还是长;后者几乎会把所有可能的算子融合方案都试一遍,显存压力更大,适合推理阶段或者 batch size 能稳定跑通之后再用。dynamic=True 那个选项其实不是省显存用的,它是告诉编译器你的输入 shape 会变化,这样它就不做死优化,反而会生成更通用的代码——代价就是编译时间更长,因为要准备多套 kernel 方案,所以如果你 batch size 固定,就别加 dynamic。我个人的经验是,对大模型微调,可以试试先不用 torch.compile 编译整个模型,而是只对 forward 里最耗时的几层(比如 attention 和 MLP)单独用 torch.compile 包一下,显存压力会小很多,提速效果也还行。另外 graph break 的问题你提了一半,我猜是编译过程中遇到不支持的操作被迫 break 导致效率下降,这其实可以通过减少 Python 层面的动态控制流(比如 if-else 或 list 操作)来缓解。你最后那个“graph br”后面是不是想说 graph break?如果是的话,可以试着把模型里的自定义算子或者数据预处理部分挪到外面,别让编译器去猜。
把dynamic=False试试,编译时间会短很多,显存占用也能降下来。
试下mode=“reduce-overhead”配合batch size再调小点,动态图模式确实编译慢但对变长输入友好。
老实讲你遇到的这个问题太典型了,我一开始玩torch.compile也是上来就炸显存。mode="reduce-overhead"其实主要是减少Python层面的调度开销,对显存优化帮助不大,真正吃显存的是编译过程中生成的中间图和临时缓存,尤其是大模型第一次编译时,TorchDynamo会做大量图捕获和算子融合,这个阶段显存峰值很容易翻倍。我自己的经验是,对于7B这种规模,先别开dynamic=True,虽然它能处理动态形状,但会让图捕获变得更复杂,编译时间和显存占用都会涨,反而得不偿失。如果你只是微调,batch size又很小,其实可以试试先不编译整个模型,只对几个计算密集的层用torch.compile,比如attention和FFN,这样能压住显存峰值。另外有个小技巧是设置torch._dynamo.config.cache_size_limit=1,限制缓存数量也能省点显存。至于graph break的问题,我猜你是不是用了某些不支持的ops,比如自定义的CUDA扩展或者控制流,这会导致图被频繁打断,编译效率暴跌,建议检查下代码里有没有类似if语句或者Python列表操作。说到底,torch.compile对大模型微调的显存友好度确实还没那么成熟,我目前还在用deepspeed的ZeRO-3配合torch.compile的partial模式,效果比全量编译稳定得多,你可以试试这条路。
说实话你这情况太真实了,我当时搞13B模型也踩过同样的坑,torch.compile在编译阶段确实会额外吃显存,尤其是默认模式会做很多graph break检查,直接给你拉爆。我后来试下来,对大模型微调最稳的反而是mode=“default”加上dynamic=False,编译时间虽然长点,但编译完跑起来显存反而比原生还省个10%左右,而且不会炸。你提到的“reduce-overhead”其实更适合小模型或者推理场景,对大模型来说它搞的算子融合反而容易让中间变量膨胀。至于dynamic=True,我理解是给动态shape用的,像你这种固定batch size的微调开了它纯属浪费编译时间,还会多存一份动态图的内存开销。另外有个坑是torch.compile跟DeepSpeed的ZeRO3搭配时会莫名多出一些临时buffer,我后来干脆关掉ZeRO3的cpu_offload才稳住。如果你只是要省显存,不如试试给torch.compile传个fullgraph=True,强制它把所有图都编译到一起,虽然第一次跑会慢到怀疑人生,但后续迭代显存和速度都稳很多。最后建议你先把一些padding操作或者动态shape的预处理放到编译外面,这种小trick往往比调mode参数管用。