最近在折腾LLM推理优化,看到大家都在吹torch.compile能白嫖30%性能,就把手里的Llama-3-8B QLoRA微调脚本改了改。结果train_step时间不降反升,显存倒是多吃了2G。查了文档说dynamic=True能处理变长输入,但我序列padding到512后shape固定了啊?还有那个mode='reduce-overhead'和默认的inductor有啥本质区别?另外给attention层加了SDPA后,编译时老报“Triton codegen failed”,回退到eager又变慢。有没有老哥在A100上实际验证过什么场景下编译收益最大?还是说我这种小batch(4)压根不适合开编译?真心求教,别让我再瞎试了。
PyTorch 2.0编译模式到底该不该开?试了几天反而更慢了
全部回复
共 63 条torch.compile这个事我踩过差不多的坑,QLoRA里那个bfloat16和gradient checkpointing跟inductor的算子融合经常打架,尤其你用了SDPA之后,Flash Attention的kernel本来就已经是手工优化的,再让Triton去codegen反而会破坏内存访问模式,报错大概率是它想融合但发现没法做向量化。我自己的经验是,小batch(4以下)加padding固定长度,编译的启动开销和显存碎片化基本会把收益吃干净,尤其A100上显存带宽那么高,eager模式的反向传播未必慢多少。你要真想试出效果,建议把batch提到8以上,或者干脆用torch.compile的fullgraph=True配合mode='max-autotune',看看CUDA graph能不能把kernel launch时间压下去。至于dynamic=True,那个是给shape真正会变的场景用的,你padding到512了shape其实是静态的,开了反而会加一层guard检查的overhead。reduce-overhead主要省的是Python和CUDA之间的同步开销,但对训练来说提升很小,推理时可能更明显。我最后是老老实实关了编译,只把SDPA保留下来,然后自己写了个简单的kernel fusion,速度反而比torch.compile默认模式快了百分之十几,这玩意儿真不是无脑开的。
说实话你这种情况我也踩过坑,小batch下torch.compile的图优化开销根本摊不薄,尤其QLoRA里那些自定义autograd函数很容易打断fusion。SDPA报错多半是Triton对某些attention mask形状支持不完善,建议先关掉编译单独跑SDPA看下真实耗时。A100上我试过,只有batch堆到16以上、序列长度固定且模型无动态控制流时,compile才有明显收益,否则纯纯负优化。你不如先试试torch._inductor.config里把conv和matmul的融合策略调激进点,或者直接换flash-attention-v2,可能比折腾编译省心。
说实话你这个问题我上周刚踩过,小batch下torch.compile的graph拆分开销远大于优化收益,尤其QLoRA里那些自定义autograd.Function基本都成编译黑盒了。SDPA报Triton codegen failed大概率是attention mask的shape没走静态路径,试试把max_seq_len写死成常量再编译。A100上我实测过,只有batch≥16且序列长度固定时,inductor才能赢回编译时间,显存多占是CUDA graph的workspace缓存,正常现象。你要是主要跑推理,不如直接换vLLM或者TensorRT-LLM,省得折腾这个。
QLoRA场景下torch.compile收益本来就玄学,微调时梯度更新和量化权重交织,inductor的图优化经常被反卷积和自定义autograd打断,白嫖性能基本是推理端的说法。你batch=4这么小,编译开销摊不薄,加上A100上SDPA本来就走flash attention内核,再套一层编译纯属重复劳动。真要吃编译红利,试试推理阶段用torch.compile+fullgraph=True,把动态shape彻底关死,batch提到16以上,收益才能出来。Triton报错大概率是attention里有residual连接触发了codegen边界,建议先关掉SDPA单独编一次看是否正常。
torch.compile这个事儿我最近也踩了不少坑,尤其是QLoRA场景下,它跟bitsandbytes的融合kernel经常打架,显存多吃2G太正常了,因为graph break之后中间张量全被保留了。你那个dynamic=True其实对固定shape没用,反而会触发recompile开销,建议直接关掉试试。SDPA报Triton codegen failed大概率是attention mask的shape没写对,或者用了bool mask,换成float mask或者把attn_mask改成None走causal路径能绕过去。reduce-overhead本质是cudnn graph + CUDA graph捕获,对小batch收益明显,但前提是模型里没有动态控制流,一旦有if或者python循环就全废了。我自己的实测是A100上batch=1、seq_len=512的Llama-2-7B生成场景,compile后decode阶段能快15%左右,但prefill反而慢10%,所以如果你想省事,还是得区分train和inference分开配。另外你提到小batch=4,我怀疑瓶颈根本不在算力而在数据加载和CPU端预处理,compile优化的是kernel launch,这个量级下反而可能被编译时间摊薄。最后建议你开TORCH_LOGS=recompiles看看有没有频繁recompile,这个日志能直接告诉你哪些地方触发了graph break,比瞎调参数靠谱得多。
我最近也在A100上试过类似的场景,torch.compile对QLoRA这种带自定义backward的优化器确实不友好,显存开销大大概率是graph捕获时把优化器状态也缓存了。SDPA那个报错可以试试把torch.backends.cuda.enable_flash_sdp设为False,有时候是flash attention的triton kernel和你的padding长度不匹配。小batch下编译启动开销占比太高了,我实测bs=8以上、序列长度固定1024时才有明显收益,你这种4的batch建议直接关掉compile专心调gradient checkpointing。另外reduce-overhead主要是减少kernel launch的python层调度,和inductor的算子融合是两条路,前者对你的场景提升很有限。
说实话你这情况我太熟了,之前拿A100试过7B和13B的SDPA编译,小batch下torch.compile的图优化收益基本被Triton kernel的启动开销吃干净了,尤其你显存还多了2G,大概率是编译时保留了中间激活做梯度回传,QLoRA的lora分支没被正确融合。dynamic=True那个参数别看文档吹得神,实际对padding后的固定shape没啥用,它主要优化的是序列长度真在变的情况,你512定长反而可能因为dynamic分支的判断逻辑拖慢速度。至于reduce-overhead和inductor,前者本质是CUDA graph捕获,能减少kernel launch但吃显存更凶,后者是默认的算子融合策略,你那个Triton报错多半是attention里某些自定义pattern没被inductor识别,建议直接试试torch.compile里加fullgraph=True或者干脆把SDPA换成flash-attn的官方实现绕开codegen。我自己的经验是,8B模型batch4这种规模,编译收益要到4090或H100上才明显,A100反而适合直接eager加混合精度,要么你试试把编译范围只包在attention层外面,别整个model都编。
torch.compile这套东西在小batch下确实容易负优化,我拿A100试过7B模型,bs4的时候显存开销和kernel launch延迟直接把收益吃光了,建议你试试把batch提到16以上或者用gradient accumulation模拟。dynamic=True对padding序列没啥用,反而会触发额外的shape分支判断拖慢速度,固定shape直接关掉就行。SDPA那个报错八成是attention mask的维度没匹配上Triton的模板,换成手动实现F.scaled_dot_product_attention传bool mask试试。我这边实际收益最大的场景是bs32+seq512的预训练,推理时用compile反而比eager慢10%左右。
小batch下编译开销根本摊不平,我试过bs=1反而掉点,你试试加大到8或16再看收益。
小batch就别折腾compile了,我试过8以下基本负优化,得把batch堆到16以上才回本。
torch.compile这个事儿吧,真不是无脑开的,我拿A100试过一堆场景,感觉收益最大的是那种大batch、固定shape、计算密集的CNN或者GPT类预训练,小batch微调反而容易负优化。你那个QLoRA本身就有lora分支和量化反量化开销,编译器很难把动态图融合干净,再加上SDPA的triton内核跟你的自定义attention不兼容,报codegen failed太正常了。我建议你先把compile关掉,直接对比一下eager模式下SDPA和普通attention的差距,大概率比编译带来的提升还明显。至于reduce-overhead,它主要是减少Python解释器的调度开销,但对显存占用是负优化,你batch才4根本没必要。dynamic=True那个坑我也踩过,它其实不是为padding设计的,是给真正变长序列用的,你shape固定了就别开,反而会触发额外的guard检查。最后说个玄学,我试过把torch._dynamo的cache_size_limit调大,有时候能减少重新编译的次数,但治标不治本。你不如先把torch.compile版本锁到2.1.0,2.2之后的triton调度策略改过,对8B这种规模的模型特别不友好。
torch.compile在小batch下确实容易负优化,我试过8的batch跑7B模型,编译开销比省下的时间还多,建议至少batch到16以上再看收益。SDPA那个报错大概率是Triton和你的CUDA版本不匹配,换个12.1的toolkit试试。另外reduce-overhead本质是cudagraphs包装,对动态shape不友好,你padding固定了shape其实可以直接开fullgraph=True,能省不少编译时间。我A100上实测过,只有序列长度和batch都翻倍时inductor才有明显优势,你这场景不如直接关掉省心。
batch太小还带QLoRA,编译开销根本摊不回来,试试batch翻倍或者直接用推理模式看收益。
小batch下编译开销本来就难摊平,试试把batch提到16再看,或者干脆关掉compile只留SDPA。
小batch+QLoRA确实容易负优化,我试过bf16+compile只在batch≥8才有提升,Triton报错建议直接换CUDA graph。
我之前也被torch.compile坑过,小batch下编译开销根本摊不平,尤其QLoRA这种带量化分支的图,inductor经常生成一堆低效kernel。你试过把batch提到8或者16吗?A100上吞吐上来后差距才明显。另外SDPA报Triton错误大概率是attention mask没走融合路径,换用xformers的memory_efficient_attention绕一下可能更稳。
说实话我之前也在A100上踩过类似的坑,torch.compile对batch size特别敏感,4这种小batch很多时候编译开销根本摊不回来,建议试试把batch提到8或16再对比。另外SDPA报Triton codegen failed大概率是attention里某些动态shape没写死,你可以试试把padding后的长度硬编码进模型forward,或者关掉dynamic=True用静态shape。至于reduce-overhead和inductor,前者主要省了Python和CUDA间的同步开销,适合小算子密集的场景,但LLM这种大算子反而可能没区别。我实际测下来,8B模型在batch=16、序列长度固定时,编译后大概能快15%左右,但前提是你得把融合算子手动调一遍,纯靠默认配置确实容易负优化。
我也踩过类似的坑,torch.compile真不是无脑开的。你那个train_step变慢其实很典型,小batch下inductor的图优化开销摊不薄,A100上batch4基本属于负优化区间,我试过batch到16以上收益才明显。dynamic=True那个参数坑更多人,它本质是让编译器为不同shape生成多个专用kernel,但你padding到512固定shape后反而触发了一堆guard检查,白白增加CPU开销,建议直接删了。
SDPA报Triton codegen failed大概率是attention里用了动态shape或者mask,编译器没法静态推断。我后来改成把mask提前构建成固定四维张量,再把非零元素用torch.where过滤掉,编译就过了。不过说实话,QLoRA场景下你大部分时间耗在LoRA的weight decomposition上,这玩意儿本身对编译不友好,收益会被吃掉大半。
如果你想验证收益,建议先用纯FP16推理模式测,不开梯度,把batch拉到32以上,seq_len固定,然后对比torch.compile的max-autotune模式。mode=reduce-overhead主要针对CUDA graph捕获,减少kernel launch次数,但小模型上graph捕获本身要额外显存,你多吃的2G大概率就是它。我最后是选择只编译attention层,其他层保持eager,收益最稳定。
我最近也在A100上试过类似的,小batch下torch.compile确实容易负优化,尤其QLoRA这种带自定义backward的,graph break一大堆,建议先关掉dynamic试试。SDPA那个报错大概率是Triton版本和CUDA不匹配,换回eager attention但保留compile有时反而稳。你要真想白嫖性能,不如把精力放在flash-attention和梯度检查点上,我这边8B模型batch=4收益比compile明显。另外reduce-overhead主要省的是CPU launch开销,但小batch下这开销本来就不大,感知不强。
小batch基本别指望compile,我试过4以下收益都是负的,试试大batch加静态shape吧。