最近在做一个基于LangChain的AI Agent项目,后端推理用的是自己微调的Llama模型,跑在PyTorch 2.0上。模型加载后推理速度有点慢,想优化一下。我查了torch.compile的新特性,据说能动态图编译加速,但又看到有人用torch.jit.script做静态图。我的场景是Agent每次调用模型时,输入长度和结构都不太一样(因为要拼接历史对话),这种情况下用compile会不会反而变慢?还是说JIT更适合动态输入?另外,如果模型里有一些自定义的注意力掩码操作,会不会影响编译效果?求有经验的大佬指点,踩坑踩得有点懵。
用PyTorch写Agent时,torch.compile和JIT到底该用哪个?
全部回复
共 145 条你这场景我太熟了,torch.compile对动态shape确实会频繁recompile,首轮慢后面才有收益,但Agent每次输入长度都变的话,加速比可能被编译开销吃掉不少。JIT对动态输入倒是友好些,但自定义注意力mask如果带Python控制流,script也容易报错,得先确保mask逻辑能完全trace。建议先试试torch.compile的mode=reduce-overhead加dynamic=True参数,配合shape padding到固定长度,实测比JIT省心。mask操作只要纯tensor运算就别怕,编译器能优化,怕的是if依赖tensor值那种分支。
torch.compile对动态输入更友好,JIT遇到变长序列容易重编译反而更慢,自定义mask用compile没问题。
说实话你这个场景我太有同感了,之前折腾过类似的动态输入Agent,最后发现torch.compile在变长序列上确实会先吃一波编译开销,但跑稳之后收益还是比JIT明显,尤其当你的batch size和长度波动不是特别离谱的时候。不过你提到自定义注意力掩码,这个得看具体实现,如果掩码是纯张量运算还好,一旦有Python控制流或者依赖输入shape的if分支,compile的图捕获可能会退化成eager模式,反而比不用还慢。我当时是先把模型里那些动态shape相关的地方尽量改成静态假设,比如padding到固定长度,再用compile,效果就上来了。JIT的话,script模式对动态shape支持更差,除非你的模型结构极其规整,否则别碰,trace模式就更别想了,历史对话一拼接就废。你不如先profile一下到底慢在模型前向还是tokenizer和预处理,很多时候Agent的瓶颈在多次前向之间的调度和显存碎片,而不是单次推理。另外如果Llama是微调过的,注意看下有没有用flash attention或者SDPA,这玩意儿配合compile的提升比单纯上JIT大得多,我甚至建议你先试试把注意力换成SDPA再决定要不要上编译。最后提醒一下,torch.compile的warmup时间在Agent这种多次小请求的场景里可能占大头,如果每次调用都间隔很久,那你得考虑用torch._dynamo的cache复用或者干脆用ONNX导出试试。
动态输入这块compile确实会反复触发重编译,输入shape一变就重新来一遍,开销可能比省下的还多。我之前也踩过类似的坑,后来用dynamic=True标记关键维度才稳住,Agent场景挺适合的。自定义掩码如果用了Python控制流大概率会graph break,建议先跑一遍explain看看断在哪。
我之前也踩过这个坑,说下我的实际感受。torch.compile对变长输入其实没那么脆,它主要是靠guard来检测形状变化,如果每次seq_len都不同,确实会频繁触发重新编译,但PyTorch 2.x有dynamic shape支持,加个dynamic=True能缓解不少。不过你这种拼接历史对话的场景,序列长度变化太频繁的话,编译开销可能把推理加速全吃回去,尤其Agent这种短请求高并发的。JIT那边静态图对动态控制流支持很一般,你自定义注意力掩码里如果有依赖运行时shape的python逻辑,script起来大概率报错或者trace成死图。我建议先别急着上compile,拿torch.profiler跑一下看瓶颈到底在kernel launch还是显存带宽,很多时候慢是因为没做KV cache或者dtype没对齐。真要试compile,先用mode="reduce-overhead"配合cudagraphs看看,但注意自定义mask操作如果涉及数据依赖的branch,编译后可能静默走eager回退,反而不如不编译。