最近在试着用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确实是硬伤,官方文档里其实提过建议用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策略固定到某个长度区间,能减少不少玄学报错。
我也遇到过这个动态shape的坑,后来发现torch.compile对变长输入确实不太友好。目前比较稳的做法是用torch.inference_mode()配合固定最大长度加mask,或者试试openxla的IE后端,对动态shape支持好一些。另外第一次过第二次挂的问题,我猜可能是torch.cuda.empty_cache()或者CUDAGraph的缓存冲突,试试关掉dynamic=True或者用mode="reduce-overhead"看看。
我也碰到过类似问题,动态shape确实是compile的痛点,特别是跨设备操作会触发重新编译。我的经验是先用固定长度跑通,再尝试torch.compile的dynamic=True参数,能缓解一部分。另外inductor后端确实有点玄学,我换成cudagraphs后端后稳定性好了一些,但速度提升没那么明显。你用的是哪个版本的PyTorch?新版本修了不少bug。
说实话这个问题我也踩过坑,torch.compile对动态shape确实不太友好,尤其是在LLM这种变长输入的场景下。你遇到的跨设备报错,我猜是模型里某些操作比如attention mask或者位置编码在编译时被静态化了,导致不同长度的输入走不同的计算路径,设备分配就乱了。我自己试下来,目前比较靠谱的做法是先用padding把输入长度固定到一个合理的最大长度,比如512或1024,然后配合torch.compile的dynamic=True参数,这样Pytorch会尽量保留一些动态维度,但依然不能保证百分百稳定。至于inductor后端抽风的问题,我也有同感,有时候第一次推理成功,第二次就报错,感觉跟算子融合时的内存分配有关,这时候可以试试换后端,比如用nvfuser或者干脆回退到eager模式。另外有个小技巧,如果你对推理速度不是极端敏感,可以先用torch.compile只编译decoder部分的前向函数,而不是整个模型,这样能减少很多莫名其妙的跨设备错误。你用的具体是什么模型架构?有些模型因为自定义算子多,兼容性会更差,可能需要改改模型代码才能跑通。
我也遇到过类似问题,动态shape确实是compile的痛点,尤其是attention里mask和cache的尺寸不一致时特别容易炸。建议试试把输入padding到固定长度,或者用torch._dynamo.config.capture_dynamic_shape=True强行开启动态支持,但稳定性还是看运气。inductor后端感觉对大模型支持还没完全成熟,我后来切回aot_eager反而稳定些,就是速度提升有限。
我也遇到过类似问题,动态shape确实坑,试下把padding固定到某个最大长度再配合dynamic=True参数试试。
说实话你遇到的这个问题我折腾过好几轮了,动态shape确实是torch.compile在LLM推理时的老冤家。即便是7B的模型,输入长度一变,内部的各种attention mask和位置编码就容易在device分配上打架,尤其是你用了padding但没对齐设备的话。我的经验是,如果实在不想固定长度,可以试试在compile时加上mode="reduce-overhead"或者dynamic=True的配置,虽然不一定完全解决,但至少能减少一些随机报错。另外,inductor后端在第一次编译时经常因为图捕获不完整导致后续挂掉,我后来换成了aot_eager或者更稳定的cudagraphs后端,虽然速度提升没那么夸张,但至少不频繁崩了。还有个偏方是手动把输入padding到某个固定长度(比如512的倍数),然后配合torch.compile的fullgraph=True,这样模型内部的操作基本不会跨设备,报错概率会低很多。不过说真的,目前PyTorch 2.0对大模型推理的编译支持还在打磨阶段,你要是追求稳定,不妨先用vLLM或者TensorRT-LLM这类专门优化过的框架,它们对动态shape的处理成熟得多。当然,如果你非要和compile死磕,建议多留意PyTorch官方github的issue区,很多人都在反馈类似的问题,说不定哪天就修好了。
试试把max_length固定,或者用dynamic参数,我上次也是这么搞定的。
动态shape确实是torch.compile的老大难,尤其是大模型推理时,padding带来的跨设备问题特别容易踩坑。我试过用torch.inference_mode()配合torch.amp.autocast,再手动把输入pad到固定长度(比如512的倍数),compile的稳定性会好很多。另外,inductor后端偶尔抽风的问题,可以试试换成aot_ts或者nvfuser,虽然速度可能慢点,但挂的概率低不少。你用的是哪个版本的PyTorch?2.0.1有个patch修过一些动态shape的bug,更新一下说不定能救。
动态shape确实容易踩坑,试试用torch._dynamo.mark_dynamic标记一下输入,或者先固定长度跑通再说。
老实说,这个问题我也踩过坑,动态shape加torch.compile确实是目前大模型推理的一个痛点。你那个跨设备报错我怀疑是某些操作在padding后产生了不同shape的中间张量,然后被JIT编译时默认分配到了不同设备上。我现在实践下来,固定输入长度是最稳的办法,比如设置几个常用的序列长度(512、1024、2048)分别编译成不同graph,推理时根据实际长度选对应的版本,虽然有点笨但至少不崩。inductor后端我也试过,第一次第二次结果不一致的问题遇到过好几次,感觉它对动态控制流的处理还不够成熟。另外你试过设置torch._dynamo.config.cache_size_limit和torch._dynamo.config.accumulated_cache_size吗?调大缓存有时候能缓解二次挂掉的玄学问题。还有一个冷门技巧是给模型加个静态的max_seq_len参数,然后用torch.compile的dynamic=False强制关闭动态shape支持,虽然会牺牲一点灵活性但稳定很多。总之目前PyTorch 2.0的compile对大模型推理还不是开箱即用,得配合一些工程trick,希望后续版本能优化这块。
动态shape确实坑,试试用torch._dynamo.config.inline_instead_of_compile过滤掉变化部分,或者先固定长度跑通再调。
动态shape确实容易踩坑,我试过把输入pad到固定长度才稳定,但会牺牲一些效率。
老实说这个问题我也踩过坑,动态shape确实是compile的硬伤,目前最稳的方案确实是把输入padding到固定长度,或者用torch._dynamo的dynamic=True参数试试,虽然会牺牲一点加速效果。inductor后端对动态图的支持确实还在完善,我第一次能跑第二次挂的情况遇到过好几次,后来换成nvfuser后端反而稳定一些。另外建议把模型里的cat和reshape操作尽量改成静态化,跨设备错误大概率是某个子模块的device没对齐。
同感,动态shape在compile下确实容易踩坑,7B模型用padding后跨设备报错我也遇到过,后来被迫把输入统一成固定长度才稳定下来。试试在torch.compile里加个dynamic=True参数,或者用backend="eager"跑一遍看看是不是模型本身有device不一致的问题?至于inductor第二次挂,我猜可能是缓存编译图时遇到了shape变化,目前感觉大模型推理还是torch.compile加固定batch/size最稳,或者直接上vLLM这类专门优化的方案省心。
刚踩过这个坑,动态shape确实是compile的硬伤,目前最稳的办法是把输入padding到固定长度,或者用torch._dynamo.allow_in_graph跳过那些动态操作。inductor玄学+1,我后来换回了默认的eager模式,虽然慢点但至少不崩。你试没试过把模型里容易跨设备的部分单独用no_grad包一下?