最近在折腾用 Llama 2 7B 做微调,看到 PyTorch 2.0 的 torch.compile 吹得很厉害,说能白嫖 30%-50% 的加速。但我试了下在单卡 A100 上跑 LoRA,torch.compile(mode="reduce-overhead") 反而比原生 torch 慢了一截,而且第一次编译要等好久。我用的 deepspeed stage2 加 bf16,是不是大模型场景下编译加速不明显,还是我哪里配错了?大佬们有没有实际对比过训练吞吐的?另外,编译后的显存开销好像也变大了,这正常吗?真心求教,不想把时间浪费在玄学调优上。
PyTorch 2.0 编译模式在大模型训练中真的比原生快很多吗?
全部回复
共 165 条实测torch.compile在小batch下反而有额外开销,大模型还是得看具体场景,建议先关掉编译跑一遍对比。
我也碰到过类似情况,小模型快,大模型配deepspeed反而负优化,感觉编译对复杂并行策略支持还不行。
说实话我也踩过类似的坑,torch.compile 在小模型或者单卡简单训练上提速确实明显,但一上大模型加 ZeRO 这种分布式策略,编译模式反而容易跟 deepspeed 的算子融合打架,尤其是第一次编译的 JIT 开销在长期训练里基本白费。显存变大也正常,因为编译会保留一些中间 buffer 来优化计算图,LoRA 这种轻量适配场景不如直接关掉编译,手动调一调 gradient checkpointing 来的实在。建议你用 profile 看一下到底是哪个 kernel 在拖后腿,很多时候瓶颈在数据加载或者通信上,跟编译关系不大。
我也在A100上试过torch.compile,LoRA场景下确实没感觉到明显加速,甚至小batch下反而变慢,感觉编译对计算密集型的大batch更友好。显存变大是正常的,因为编译会生成额外的缓存和优化过的算子。建议试试mode=“default”或者关掉一些图优化看看,有时候deepspeed的overlap和编译会冲突。另外第一次编译慢是通病,第二次跑会好很多,但微调这种动态图结构确实收益不大。
我自己的测试结果跟你差不多,torch.compile 在小 batch 或者简单模型上确实有提升,但一上大模型加分布式训练,那点加速经常被编译开销和显存占用吃掉,尤其第一次编译卡很久,实际跑起来反而得不偿失。我后来直接关掉 compile 换 full bf16 + flash attention,吞吐反而更稳。你显存变大是正常的,因为编译会保留一些中间图和优化缓存,建议你试试 mode=“default” 或者干脆只在推理时开 compile。
你遇到的这个情况挺正常的,torch.compile在大模型场景下收益确实没那么稳定,尤其是和Deepspeed混用时,编译开销经常覆盖掉加速。我试过类似配置,发现用mode="max-autotune"反而更慢,但改成"default"模式后训练吞吐有5%-10%的提升。显存变大也是老问题了,编译图会多占些缓存,建议关掉一些fusion选项试试。真要白嫖加速,不如先检查下CUDA graph和flash attention的兼容性。
我最近也在折腾类似的东西,torch.compile 在小模型或者单卡推理上确实有提升,但大模型训练场景下收益真没那么神,尤其是和 DeepSpeed 混用的时候经常出兼容性问题。你说的首次编译慢和显存上涨我这边也遇到了,感觉官方吹的 30% 加速更多是理想情况下的算子融合结果,实际 LoRA 微调里瓶颈可能根本不在计算上。建议你试下 mode="max-autotune" 或者干脆关掉编译先跑个 baseline,有时候减少显存碎片反而更稳。
你说的reduce-overhead确实容易翻车,我换mode="max-autotune"后才勉强追上原生速度,编译开销在7B上挺玄学的。
torch.compile在小模型上提速明显,大模型尤其带deepspeed时确实容易负优化,建议先关掉编译试试纯原生对比。
老实说我和你情况差不多,在单卡上试torch.compile跑微调,收益真的不大,尤其是第一次编译那段时间太折磨了。我后来看了一些讨论,感觉compile对计算密集型的大batch和大模型更友好,像LoRA这种小改动反而容易因为图编译和显存碎片拖慢速度。你显存变大也正常,因为编译会保留一些辅助结构,建议试试mode="default"或者关掉deepspeed的某些优化看看有没有变化。
同款配置踩过类似的坑,torch.compile在LoRA这种小参数量微调场景下收益确实不大,尤其开reduce-overhead后显存占用反而涨了5%-10%,第一次编译那几分钟更是煎熬。不过换成fullgraph模式配合deepspeed stage3,我们实测7B全量微调大概能提15%-20%的吞吐,但前提是把动态shape和Python控制流都固化掉。你试过关闭deepspeed的offload再对比吗?有时候编译和ZeRO的显存管理会打架。
同感,torch.compile在小batch和大模型上真的容易翻车,尤其配合deepspeed stage2时,显存和编译开销都会涨。我试过把mode改成default或者max-autotune,反而比reduce-overhead稳定点,但第一次编译时间确实劝退。感觉编译对计算密集型场景(比如大卷积或大batch推理)更友好,LoRA这种小参数量微调反而容易被调度和优化器拖累。建议试试先不用deepspeed,纯torch.compile跑几步看看基线,另外检查下是否开了动态shape,那个也会让编译缓存失效。
老实说我也踩过类似的坑,torch.compile在小batch size下反而有负优化,大模型场景尤其明显。你试的reduce-overhead模式对动态图和复杂控制流不太友好,换成mode="max-autotune"或者干脆默认模式可能更好些,不过第一次编译确实要等很久。显存变大也正常,因为编译会做算子融合和张量布局优化,临时占用会高一些。我自己的经验是batch size堆到够大、序列长度固定时编译收益才明显,LoRA这类轻量微调反而容易亏。
说实话我也踩过类似的坑,torch.compile在纯Eager模式下小模型收益明显,但一旦叠上Deepspeed和LoRA这种外部封装,图优化反而容易跟梯度切分打架。你试试把mode换成“max-autotune”或者干脆关掉dynamic shape,有时候reduce-overhead在bf16下会额外插入转换节点。显存变大挺正常的,编译会保留一些中间buffer,不过如果超过5%就得查查是不是跟activation checkpointing冲突了。对了,你用的是torch 2.1+吗?老版本对Llama的attention优化支持很差。
说实话你这个问题我太有共鸣了,之前我在A100上跑LLaMA微调也踩过一模一样的坑。torch.compile在CV或者小模型上确实能带来惊喜,但大模型场景下,尤其是和deepspeed、bf16混在一起时,编译的收益经常被通信和显存管理的开销抵消掉。我试过几次,reduce-overhead模式反而会让显存峰值涨不少,因为编译器会做算子融合和额外的buffer分配,这在长序列训练里特别明显。我的经验是,如果你已经用了deepspeed stage2,那编译带来的计算优化空间其实很有限,因为瓶颈更多在带宽和显存访问上,而不是算子执行效率。另外第一次编译那几十分钟确实蛋疼,但如果你会用torch.compile的dynamic=True或者先跑个小batch预热,能稍微缓解一点。我后来干脆放弃编译,改成手动调gradient checkpointing和batch size,反而吞吐更稳定。你检查下是不是因为LoRA的low-rank分支太轻量,编译反而增加了调度开销?我觉得你可以试试不加deepspeed,纯用FSDP加compile对比一下,说不定结论会反转。
跟你同配置,A100加deepspeed stage2,实测compile在LoRA上基本没收益,甚至小batch下还倒吸。后来发现torch.compile对动态shape和显存优化那块特别敏感,你把max_length固定死,再试试mode="max-autotune",吞吐能上来一点,但显存确实会多吃几个G,这玩意儿更适合大batch纯推理或者全参数训练,微调场景真没必要硬上。
这情况太真实了,reduce-overhead在deepspeed下编译收益本来就容易被通信吃掉,显存涨是图优化缓存导致的,正常。
小模型loRA吃不满A100,编译那点优化还不够抵消额外开销的,换全参微调或者加大batch再试试。
说实话我跟你遇到的情况几乎一模一样,单卡A100上跑LoRA,torch.compile开reduce-overhead反而慢,后来我干脆关掉换回原生训练了。后来仔细翻了翻issue,发现这玩意在7B这种规模下,编译优化主要吃的是算子融合和kernel自动调优,但LoRA本身瓶颈在显存带宽和通信上,编译能省的反而被CUDA graph的捕获和重放开销抵消了。你试过mode="default"或者"max-autotune"吗?reduce-overhead那个模式我实测对动态shape特别敏感,Llama的attention里如果带了position encoding的变长mask,很容易触发重新编译,那延迟比省下来的还多。另外显存变大是正常的,因为编译会预留一些中间buffer和graph内存,尤其max-autotune模式下能多占2-3G,这在7B上挺伤的。我后来换了个思路,用deepspeed的zero stage3加flash-attention,不开compile,吞吐反而比开了compile的stage2高15%左右。建议你可以试试把bf16换成fp16,有时候A100上bf16的编译路径反而没有fp16优化得透彻。最后想问下你的batch size设了多少?我怀疑你那个场景下小batch让编译的固定开销占比太大了。
我也踩过类似的坑,先说结论:torch.compile在大模型训练里真不是无脑白嫖,尤其你还在用deepspeed stage2。reduce-overhead模式本身是为小batch推理优化的,训练场景下它那套CUDA graph捕获反而会跟deepspeed的梯度分区和参数更新产生冲突,吞吐掉是正常的。我试过在7B LoRA上切mode="max-autotune",配合fullgraph=True,勉强能跟原生持平,但编译时间够我刷两集剧了。显存增加大概率是因为编译缓存了额外的工作区,加上graph捕获的静态buffer,这在长序列下会更明显,你可以试试把max-autotune的cache_size调小或者关掉triton的cudagraph。个人建议如果你不是特别在意那20%的极限吞吐,现阶段用原生torch加deepspeed的offload更省心,毕竟编译炸一次OOM排查起来真能让人头秃。另外看下你是不是没设torch._dynamo.config.capture_scalar_outputs,有些算子因为标量返回被回退到eager,反而更慢。
小模型上编译提速明显,7B这种规模反而容易吃编译开销,显存变大也正常,建议直接关掉对比试试。
照你这配置,reduce-overhead本来就不适合训练,试试默认模式或者干脆别开,省下的时间多调调lr更实在。