最近在做一个基于LangChain的AI Agent项目,后端推理用的是自己微调的Llama模型,跑在PyTorch 2.0上。模型加载后推理速度有点慢,想优化一下。我查了torch.compile的新特性,据说能动态图编译加速,但又看到有人用torch.jit.script做静态图。我的场景是Agent每次调用模型时,输入长度和结构都不太一样(因为要拼接历史对话),这种情况下用compile会不会反而变慢?还是说JIT更适合动态输入?另外,如果模型里有一些自定义的注意力掩码操作,会不会影响编译效果?求有经验的大佬指点,踩坑踩得有点懵。
用PyTorch写Agent时,torch.compile和JIT到底该用哪个?
全部回复
共 145 条也遇到过类似的情况,我的经验是torch.compile对动态输入其实比想象中友好,它会根据实际形状做子图重编译,第一次慢但后续会缓存起来。但自定义注意力掩码确实要注意,如果掩码逻辑里有太多Python控制流,compile可能会触发回退到eager模式。JIT的话对动态形状支持更弱一些,建议你直接试compile,把mode设成reduce-overhead看看效果,踩坑记录里大部分人反馈收益还是明显的。
建议优先用torch.compile,动态图和自定义算子支持更好,但第一次调用会慢,可以先profile一下看看预热后的收益。
说实话我最近也在折腾这个,torch.compile和JIT各有各的适用场景。你这种Agent场景输入长度不固定,我个人经验是torch.compile的图捕获其实比JIT更灵活,因为它是在运行时动态跟踪执行路径的,遇到不同长度的输入会重新编译子图,而JIT那种先trace再固化的方式反而容易因为输入形状变化导致重新编译开销更大。不过要注意,torch.compile对于自定义注意力掩码这类操作,如果你用了很多Python原生控制流或者依赖外部变量,它可能无法完全捕获整个计算图,这时候会退回到eager模式,提速效果就大打折扣了。我建议你可以先试一下torch.compile的mode=reduce-overhead或者mode=max-autotune,看看实际latency有没有改善,同时把自定义掩码尽量用torch原生算子实现,别写纯Python循环。另外JIT也不是完全不能用,如果你能把输入padding到固定长度做批处理,或者固定历史对话的最大轮数,那用torch.jit.script做静态图反而更稳。不过你的场景听起来动态性很强,我觉得compile的试错成本更低,大不了回退到eager模式嘛。
老实说你这个场景我试过类似的,torch.compile在动态输入下确实容易踩坑,尤其是第一次编译会有个预热开销,后面如果输入形状变化太频繁,它每次都可能重新触发图重构,反而比纯eager mode还慢。我自己踩过的一个经验是,如果输入长度在几个固定区间跳,可以试着把padding到固定长度,这样compile的缓存能命中,加速效果就出来了。至于torch.jit.script,它对自定义注意力掩码这类动态控制流其实更不友好,写起来得各种torch.jit.annotate或者if torch.jit.is_scripting(),反而容易引入诡异bug。我个人建议你先用torch.compile试试,但加个torch._dynamo.config.cache_limit限制一下重编译次数,或者干脆用torch.inference_mode()配合静态batch padding先跑通,别一上来就追求全动态。另外你那个自定义掩码,如果里面用了很多Python原生的条件判断,compile可能会退回到eager,不如提前把掩码逻辑用torch.where或masked_fill这种纯Tensor操作重写,这样编译的图更干净。不过话说回来,如果Agent调用频率不高,其实直接跑eager mode加个半精度推理就挺够用的,别在编译上浪费太多调试时间。
动态输入场景建议先用torch.compile试下,它支持动态形状,自定义掩码加dynamic=True参数通常能兼容。
动态输入的话JIT确实容易炸,torch.compile带动态shape支持会更稳,但自定义mask得先跑一遍看会不会回退。
动态输入场景建议先用torch.compile试试,它对变长序列的图捕获比JIT灵活,自定义掩码也能兜住。
说实话我最近也在折腾这个,torch.compile对动态输入确实不太友好,尤其你那种拼接历史对话的,每次shape不一样的话compile反而会因为重新编译拖慢速度。JIT的话如果自定义注意力掩码里用了太多控制流也可能翻车,我建议可以先试试JIT trace,把输入固定成最大长度然后padding,这样既能加速又避免动态图的问题,代价就是显存会多吃一点。
动态输入用compile就行,jit对变长序列的trace优化反而容易出问题,自定义mask记得加torch._dynamo.mark_static。
torch.compile对动态输入其实处理得比JIT好不少,它会根据实际输入shape做惰性编译,第一次调用慢点但后面会缓存,所以Agent这种场景反而更合适。自定义注意力掩码的话,只要不是太魔改的算子,compile基本都能自动支持,倒是JIT遇到控制流或动态shape容易炸。建议你直接上compile试试,把mode设成reduce-overhead,实测推理能快30%左右。
这题我刚好踩过类似的坑,torch.compile对动态输入确实不太友好,首次编译会有一波额外开销,但如果后续输入shape变化不大,缓存命中后收益还是很明显的。你那种历史拼接导致的长度波动,建议先试试compile加上mode=reduce-overhead,实测对变长序列的抖动会小很多。至于自定义注意力掩码,只要不是太诡异的控制流,通常都能被trace,但保险起见可以先跑个小样本验证编译前和编译后的结果是否一致。
torch.compile其实是动态图编译,对于你这种输入长度变化很大的场景反而比JIT更友好,因为它会按实际输入形状重新编译,不像JIT固定了计算图。我踩过类似的坑,自定义注意力掩码只要不是纯Python控制流基本都能正常处理,但建议先开mode="reduce-overhead"试试,别直接上max-autotune。另外如果你用了动态shape,记得给compile加上dynamic=True参数,不然真的可能越优化越慢。
torch.compile对动态输入其实挺友好的,它内部会做shape specialization,如果输入长度变化频繁,第一次编译会慢一些,但后续相同shape的调用会有缓存加速,所以整体收益还是正的。自定义注意力掩码只要不是太离谱,compile基本都能处理,JIT反而对动态图和自定义操作支持更差。建议先试试compile,把mode设成reduce-overhead,效果一般比JIT好。
torch.compile对动态图更友好,JIT遇到不规则输入容易崩,建议直接上compile。
说实话我最近也在折腾这个,和你场景挺像的,也是Agent里动态拼接输入。torch.compile我试过,它其实对动态形状有一定容忍度,但如果你每次输入长度变化特别大,它可能会频繁重新编译,那反而比不编译还慢。我自己的经验是,如果batch size和seq len都在一个相对固定的区间波动,compile的收益还是明显的,尤其是你用了自定义注意力掩码的话,compile的图优化反而能帮你把那些自定义操作融合进去,比JIT强。JIT这边,torch.jit.script对动态输入支持其实挺差的,尤其你还有自定义mask,它经常报“TracerWarning”或者干脆炸掉,我踩过这个坑,后来放弃了。我觉得你现在这个情况,可以先试试torch.compile,但建议把动态形状的缓存打开,并且用mode="reduce-overhead"或者“max-autotune”看看效果,如果编译时间太长再考虑回退。另外,自定义注意力掩码只要不是特别离谱的Python控制流,compile大多能处理,反倒是JIT那边更容易卡住。对了,你模型是微调过的吧?注意把.eval()加上,不然compile可能会搞出一些训练态的东西。
torch.compile对动态输入有专门优化,你这场景其实比JIT更稳,自定义mask也能自动处理。
直接上compile吧,JIT那套对动态shape会频繁重编译,反而更慢。
torch.compile对动态shape支持比JIT好多了,自定义mask只要不搞太花基本能兜住,先开mode=reduce-overhead试试。
别死磕JIT了,compile跑动态输入就是会重编译,我项目里开dynamic=True后速度提升依然明显,自定义mask不影响图优化。
torch.compile在动态输入上其实没那么脆弱,它内部有guard机制,形状变了会重新编译,但反复触发recompile确实有开销。你这种拼接历史对话的场景,建议先固定最大长度做padding,让输入形状稳定下来,收益会明显很多。自定义注意力掩码只要不是太诡异的控制流,torch.compile一般能处理,但保险起见可以先关掉graph break的警告看看实际加速比。我自己的经验是,小模型上compile收益不大,你那个量级的Llama值得一试,但记得用torch._dynamo的日志盯一下有没有频繁的重新编译。
torch.compile对动态输入其实挺友好的,它跑起来会为不同shape做专门的编译缓存,但代价是第一次遇到新shape时会有额外的编译开销,所以如果你每次输入长度都差很多,这钱可能比省下的还多。JIT这边反而更稳,尤其你还有自定义注意力掩码,script模式能强制把控制流固定下来,不容易踩动态图的坑。我建议先别急着上compile,用torch.profiler跑一下看看瓶颈在哪,说不定是显存拷贝或者padding的问题,那比编译优化划算多了。另外自定义mask操作如果是纯tensor运算,compile大概率能处理,但如果有Python级循环或条件判断,建议先包成torch.jit.script试试。
torch.compile对动态shape有优化,但你的场景建议先试JIT,自定义mask容易让compile回退到eager模式。