最近在做一个基于LangChain的AI Agent项目,后端推理用的是自己微调的Llama模型,跑在PyTorch 2.0上。模型加载后推理速度有点慢,想优化一下。我查了torch.compile的新特性,据说能动态图编译加速,但又看到有人用torch.jit.script做静态图。我的场景是Agent每次调用模型时,输入长度和结构都不太一样(因为要拼接历史对话),这种情况下用compile会不会反而变慢?还是说JIT更适合动态输入?另外,如果模型里有一些自定义的注意力掩码操作,会不会影响编译效果?求有经验的大佬指点,踩坑踩得有点懵。
楼主
22小时前
用PyTorch写Agent时,torch.compile和JIT到底该用哪个?
请 登录 后发表回复
全部回复
共 2 条
2楼
6小时前
torch.compile对动态图支持更好,你这种变长输入场景建议优先试它,自定义掩码也能自动处理。
3楼
1小时前
torch.compile 对动态输入其实挺友好的,尤其是 PyTorch 2.0 以后,它会在运行时根据实际输入形状做即时编译,不像 JIT 那样需要你提前固定结构。不过你那个自定义注意力掩码确实是个坑,有的操作比如 torch.where 或者某些条件分支可能会触发回退到 eager 模式,建议先给 compile 开个 mode="reduce-overhead" 试试,同时留意下编译日志里有没有“fallback”字样。要是实在折腾不动,可以考虑先用 JIT 把模型里输入无关的部分 trace 成静态图,剩下动态部分留着跑 eager,算是个折中方案。