最近在试着用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推理时总报错,是我打开方式不对吗?
全部回复
共 5 条老实说我也被torch.compile折磨过,动态shape确实是硬伤,官方文档里其实提过建议用torch._dynamo.config.capture_dynamic_shape=True,但实测对LLM这类模型效果还是不太稳。我个人的做法是固定最大长度+动态padding到那个值,虽然浪费点显存但至少不崩了。inductor后端在第一次推理时会做编译缓存,第二次挂可能是缓存命中后设备分配出了问题,建议试试mode="reduce-overhead",或者干脆回退到eager模式跑推理,编译留给训练阶段。
这问题我也踩过坑,动态shape确实是compile的硬伤,尤其是LLM推理时padding会导致跨设备报错,我后来是强制固定输入长度+静态mask才勉强跑通。inductor后端稳定性确实有点玄学,建议试试torch.compile(dynamic=False)或者用图的capture模式,虽然牺牲一点灵活性但至少不挂。另外你模型里如果有自定义的attention实现,可能得手动标记一下device,不然第一次能过第二次就崩很常见。
老实说,你这个坑我也踩过,而且踩得挺深。7B模型用torch.compile确实对动态shape非常敏感,特别是注意力层那些需要计算长度的操作,一旦padding后的mask或者position id和设备不统一,就容易报那个跨设备错误。我后来发现一个比较省心的做法是先用torch.jit.script把模型里一些稳定的小模块单独trace掉,再对剩下部分用compile,虽然有点麻烦但至少能跑起来。至于你提到的inductor后端不稳定,我体感上也是时好时坏,有时候换用cudagraphs或者nvfuser反而更稳,但不同模型表现差异很大。我自己试下来,如果输入长度变化确实大,不如直接放弃compile,改用flash attention加torch.inference_mode,推理速度其实差不了太多,还省得debug到怀疑人生。另外你检查过模型里有没有条件分支或者动态循环吗?那种结构compile经常会二次编译导致第一次能跑第二次就挂,我遇到过好几次。总之这玩意还在快速迭代中,现在稳定落地大模型推理还是得靠triton或者vLLM那种专门优化过的方案,个人折腾成本确实高。
我也遇到过类似问题,动态shape确实是torch.compile的痛点,特别是带padding的输入容易触发重编译或设备不一致。目前我这边比较稳的做法是先固定输入长度到某个上限,用max_length+attention mask去处理,这样compile的图能稳定下来。另外inductor后端确实有点玄学,我换回默认的dynamo后端反而更稳一些,你可以试试看。还有个小技巧是加上mode="reduce-overhead"或者设置dynamic=False,能减少一些莫名其妙的报错。
试试把dynamic=True参数加到torch.compile里,或者用padding策略固定到某个长度区间,能减少不少玄学报错。