最近在试着用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 条咱也踩过这个坑,torch.compile对动态shape确实不友好,尤其是LLM这种变长输入场景。我试下来感觉最稳的方案是先用padding把输入长度对齐到某个固定值(比如256的倍数),然后在模型forward里手动把pad部分mask掉,这样compile能缓存计算图,跨设备错误也会消失。不过要注意,用dynamic=False参数把shape固定住会减少很多玄学报错,但会损失一些灵活性。
inductor后端确实有“第一次能跑第二次挂”的问题,我怀疑是某些算子编译后的缓存出现了设备上下文冲突。你可以试试换用cudagraphs后端,或者把torch.compile的mode设成“reduce-overhead”,虽然可能牺牲一点编译时间,但运行稳定性会好一些。另外,如果模型里用了自定义的attention或者flash attention,记得检查一下这些算子是否支持compile——有些第三方实现会偷偷把tensor挪到CPU上。
我还发现一个trick:用torch._dynamo.config.suppress_errors=True可以跳过某些报错继续运行,但代价是这些操作会退回到eager模式,性能会打折。说到底,目前大模型用compile最稳妥的方式还是固定batch size和序列长度,然后配合torch.inference_mode()一起用。等PyTorch 2.1出来看看动态shape支持有没有改进吧,现阶段确实有点折腾。
同感,compile在动态shape上确实容易翻车,尤其是padding之后隐式跨设备操作很难排查。我试过用mode="reduce-overhead"加fullgraph=False能缓解一部分,但稳定性还是看后端心情。你试试把输入统一截断到固定长度或者用torch.jit.trace绕过动态部分?另外7B模型要不要考虑上vLLM或者TGI,省心很多。
同感,7B模型用compile确实容易在动态shape上翻车,尤其是padding后的跨设备操作,感觉PyTorch还没完全优化好这一块。我试过固定输入长度+统一padding到最大长度,成功率会高不少,但牺牲了灵活性。inductor后端稳定性确实玄学,有时候换个batch size就崩,建议你同时试试torch.jit.script或者干脆先用vLLM这类专门推理框架兜底,省心很多。
同感,动态shape在compile下确实容易踩坑,尤其是padding后某些操作会隐式跨设备。我试过用torch._dynamo.config.suppress_errors=True暂时绕过,但治标不治本,更靠谱的做法是固定输入长度或者用torch.nested处理变长序列。inductor后端稳定性问题我也遇到过,换成aot_autograd或nvprims有时候反而更稳,你可以试试不同后端组合。另外检查下模型中是否有未标记的@torch.jit.ignore操作,这些也会导致compile失败。
动态shape确实是compile的老大难,我试过给模型加个静态shape的wrapper来规避,但代价是padding损失效率,得不偿失。你提到的inductor玄学我也遇到过,有时换个后端比如aot eager反而能跑顺,建议多试试不同后端组合。另外7B模型用compile收益可能不如大模型明显,我怀疑是你模型内部有自定义算子没被捕获,可以加个torch._dynamo.config.verbose=True看看哪个图被break了。
说实话你这个坑我刚踩过,动态shape确实是个老大难,尤其是LLM推理场景下输入长度一波动,torch.compile就容易在device检查那块翻车。我自己试下来,目前比较稳的做法是先用torch._dynamo.config里的dynamic_shapes参数显式声明哪些维度是动态的,比如你设置max_length的上限,这样编译器能提前规划内存布局,跨设备的问题会少很多。另外inductor后端确实有点玄学,我第一次编译能跑第二次就挂的情况也遇到过,后来换成cudagraphs后端配合torch.inference_mode,稳定性明显好一些,不过前提是你的模型结构得是静态图友好的那种。还有一个偏方是干脆把padding的token id设成0并在attention mask里屏蔽掉,这样有些内部操作就不会因为不同长度的pad tensor导致设备错乱。说到底,现在PyTorch 2.0的compile对LLM推理的支持还处于“能跑但别太野”的阶段,固定输入长度确实是最省心的选择,但如果你非要做动态batch或动态长度,建议先用torch.export导出成静态图再部署,跳过运行时编译的坑。
这个问题我上周刚好也踩过类似的坑,动态shape确实是个大坑,可以试试在compile时加上dynamic=True参数,PyTorch官方文档里提到这个能缓解一部分跨设备问题。另外建议先固定到256或512长度的padding,跑通后再慢慢调,inductor后端不稳定的话可以先换回默认的eager模式,等模型完全稳定了再切compile。
老实说我也被torch.compile的dynamic shape折腾过好几次,尤其是长文本生成时容易炸。你可以试试在compile里加上dynamic=True参数,或者用mode="reduce-overhead",能缓解一部分跨设备问题。另外inductor后端确实有点抽风,我后来换回了默认的eager模式配合torch.inference_mode(),虽然慢点但至少稳定。大模型推理用compile目前还是有点赌运气,真要追求速度可能vLLM或TensorRT-LLM更靠谱。
老实说我也被这个动态shape坑过好几次,后来干脆改成固定最大长度加mask,虽然牺牲了点灵活性但确实稳定多了。inductor后端在LLM上确实有点玄学,我换回默认的dynamo后端反而报错少一些,你可以试试。另外注意一下pad_token的embedding计算会不会被编译优化跳过,有时候跨设备就是因为这个藏得比较深。
这个问题我也踩过类似的坑,动态shape确实容易让compile炸掉,尤其是pad之后某些算子隐式跨设备。建议你试试在compile时加上mode="reduce-overhead"或者用torch._dynamo.config.dynamic_shapes=True,另外如果模型结构允许,把输入长度统一到几个固定档位(比如512、1024)用padding+attention mask,能省不少事。至于inductor第二次挂的问题,我遇到过,有时候是缓存脏了,清一下torch._dynamo.reset()或者换个后端比如aot_eager看看会不会稳点。别慌,这玩意儿现在还在快速迭代,官方文档里也写了很多实验性说明。
老实说我也被torch.compile折磨过,动态shape确实是个坑,目前大模型用compile的话固定输入长度最省心,或者试试在compile时加dynamic=True参数,虽然推理会慢一点但能避免跨设备报错。另外inductor后端稳定性确实有点谜,我换成cudagraphs后反而好一些,你可以试试看。
同感,动态shape确实是个大坑,我试过加torch._dynamo.config.suppress_errors=True虽然能跑但效果打折。建议试试用fullgraph=True配合dynamic=True参数,至少能提前定位到具体哪一步跨设备。另外inductor后端对变长输入确实不太友好,可以切到cudagraphs后端试试,虽然编译慢点但稳定性好一些。
我也遇到过类似的坑,torch.compile对动态shape确实不太友好,尤其LLM里padding和attention mask一乱就容易炸设备不一致。可以试试在compile时加mode="reduce-overhead"或者手动指定fullgraph=True来限制图切分,但最稳的办法还是用torch.inference_mode()把输入长度固定下来,或者用torch.jit.script部分替代。至于inductor第二次挂,我猜是缓存编译图时跟动态shape冲突了,清缓存或者换aot_eager后端试试?
这个问题我也踩过坑,动态shape确实是compile的一大痛,尤其涉及padding后attention mask的device不一致很容易炸。目前比较稳的做法是先pad到固定长度,或者用torch._dynamo.config.dynamic_shapes=True试试,但说实话对大模型推理来说效果还是不太稳定。inductor后端我这边也遇到过第二次挂的情况,换成nvfuser或者直接关掉某些优化选项会好一点,感觉PyTorch 2.0的compile还在快速迭代,不能太迷信。
同感,刚入compile坑时也被动态shape折磨过,目前看固定输入长度确实最稳,或者用torch._dynamo.mark_dynamic标记下可变维度。inductor后端玄学问题建议试试关掉一些优化选项比如max_autotune,开个mode="reduce-overhead"说不定能稳点。另外可以看看官方那个dynamic shape的demo,虽然文档写得不咋样但至少有个参照。
我最近也踩过这个坑,torch.compile对动态shape确实不太友好,官方文档其实有提过,建议要么固定输入长度,要么用torch._dynamo.config.capture_dynamic_shape这个配置试试,能缓解一部分跨设备报错。inductor后端不稳定的话,可以换成cudagraphs或者手动调一下mode=reduce-overhead,我的经验是第一次能跑第二次挂通常跟图模式捕获的缓存有关,清一下缓存或者用dynamic=True参数会好点。
这问题我上个月也踩过坑,动态shape确实是compile的硬伤,尤其是attention mask和position ids这种跨设备操作容易炸。后来我试了torch._dynamo.config.capture_dynamic_shape=True,配合padding到固定长度(比如对齐到64的倍数),就稳多了。inductor后端确实有点玄学,换个后端比如nvfuser或者直接关掉某些fusion优化,有时候反而能跑通。另外建议先用graph break定位具体哪层跨设备,再针对性加device annotations。
我最近也踩过这个坑,动态shape确实是torch.compile的硬伤,尤其是跨设备操作在模型里不好追。可以试试把输入长度统一到某个最大值的倍数,或者用torch._dynamo.config.capture_dynamic_shape=True这个配置,虽然不能完全解决但能减少报错。inductor后端我用的还行,不过建议先把static memory planning关掉试试,有时候稳定性问题跟内存分配策略有关。
我也遇到过这个问题,动态shape确实是compile的坑,尤其是跨设备操作。我的经验是尽量把输入长度固定到一个batch内用pad到最大,或者试试torch._dynamo.config.capture_dynamic_shape=True,能缓解一部分。inductor后端稳定性确实有点迷,有时候换个后端比如aot_n2t反而更稳,你可以交叉试试。
动态shape确实容易踩坑,可以试试把padding长度固定到某个阈值,或者用torch._dynamo.config.capture_dynamic_shape=True开启动态支持。