最近在试着用PyTorch 2.0的torch.compile优化一个7B的对话模型推理速度,但一跑就报“RuntimeError: Expected all tensors to be on the same device”,debug了半天发现是动态shape的问题。我的输入长度变化比较大,用padding后模型内部有些操作还是跨设备了。想问问大家,现在大模型用compile的最佳实践是什么?是不是必须固定输入长度,或者有什么配置能避免这种报错?另外,我试了用inductor后端,但有时候模型第一次推理能过,第二次就挂,感觉稳定性还是有点玄学。真诚求教,别笑我菜。
PyTorch 2.0 compile在LLM推理时总报错,是我打开方式不对吗?
全部回复
共 182 条动态shape确实是torch.compile踩坑的重灾区,你遇到的设备不一致报错往往不是真的跨设备了,而是guard失败后图被重新捕获时,某些中间tensor的device信息丢了,报错信息有点误导人。我自己的经验是7B这个量级别一上来就全图compile,先把model拆开,只对transformer block或者attention层做区域编译,效果反而更稳。固定输入长度不是必须的,但你要么用mark_only_dynamic加dynamic=True显式标记动态维度,要么就接受它每次遇到新shape都重新编译的代价。inductor第二次挂掉大概率是cache命中逻辑和动态shape的符号推导打架,可以试试开mode="reduce-overhead"之外再配fullgraph=False,让它允许图断裂。另外你padding之后跨设备,检查一下position_ids或者attention_mask是不是在某个分支里被重新创建到了cpu上,这种细节特别容易漏。说实话现在compile在大模型推理上还没到开箱即用的程度,社区里很多人都是半用半关,关键路径手写融合更靠谱。
动态shape确实是torch.compile踩坑最多的地方,你遇到的跨设备报错大概率不是真的设备不一致,而是guard失败后recompile时某些graph break把tensor状态搞乱了,报错信息有点误导人。我自己的经验是先用mode="reduce-overhead"配合dynamic=True试一下,但LLM里RoPE和KV cache那几段对shape特别敏感,经常第一次跑通第二次就崩,基本就是cache命中逻辑和编译缓存打架。固定输入长度确实最稳,但对话场景不现实,可以试试把prefill和decode分开编译,decode阶段尽量用static cache。inductor后端不稳定这事不怪你,7B模型上graph break太多的话,编译出来的东西还不如eager快。建议先开TORCH_LOGS=graph_breaks看看哪里断了,很多时候是某个自定义op或者mask操作在作妖。