最近在试着用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在LLM推理上的一个老大难。你试过用torch._dynamo.config里的dynamic_shapes参数吗?把那个打开然后显式指定输入维度的动态范围,有些场景能缓解设备不一致的报错。不过说实话,7B模型跑compile,inductor后端的高频调度和CUDA graph缓存确实容易在变长输入时抽风,我自己的经验是先用torch.inference_mode()配合静态输入跑通一次,再慢慢加padding的mask处理逻辑,但稳定性依然看运气。另外,如果你是用decoder-only架构,可以试试把attention的mask用fixed shape的因果掩码代替动态计算,这样compile更容易捕获计算图。不过话说回来,现在社区里大厂的做法好像还是倾向于用vLLM或者TensorRT-LLM这类专用框架,PyTorch的原生compile在推理场景下确实还没到无脑用的地步,你遇到这些坑挺正常的,不是你的问题。
同感,动态shape下compile确实容易翻车,试试用torch._dynamo.mark_dynamic标记一下可能会稳点。
这问题我踩过一模一样的坑,动态shape在compile下确实容易触发隐式设备同步,尤其attention mask或者position id那块。建议先试试torch._dynamo.mark_dynamic把变化的维度显式标出来,或者直接用静态shape跑一遍验证逻辑,能过再往上加动态。另外inductor第二次挂大概率是graph缓存和cudnn benchmark的玄学冲突,可以试试torch._inductor.config.triton.cudnn_benchmark=False,或者干脆用reduce_overhead模式,虽然慢点但稳定。你7B模型要是显存不紧,不如直接固定最大长度padding,省心很多。
这问题我也踩过,动态shape确实是torch.compile的痛点,尤其llama这种带cache的模型。我试过给输入长度设个上限然后padding到固定值,虽然浪费点显存但至少稳定不报错。inductor后端我也有过一模一样的“第一次过第二次挂”的情况,后来发现是cudagraph和动态shape冲突,关掉cudagraphs或者用mode=“reduce-overhead”反而好点。不过说实话,7B模型直接上vLLM或者TensorRT-LLM可能更省心,compile这功能更适合静态batch的生产场景。
这问题我熟,之前也踩过同样的坑。动态shape在compile下确实容易触发设备检查的误报,特别是padding后某些算子内部逻辑没跟上。建议试试把输入长度固定到几个档位,比如128/256/512,用bucket方式做静态化,比硬padding稳很多。另外inductor第一次能跑第二次挂大概率是缓存复用的问题,可以试试设置torch._dynamo.config.cache_size_limit=0或者换成cudagraphs后端,虽然慢点但至少不玄学。
试过把padding改成右对齐+固定max_len,compile稳多了,动态shape确实容易触发设备断言。
动态shape这块确实是torch.compile的老大难,尤其是LLM推理场景,我建议先把max_seq_len固定下来,配合padding_mask一起传给模型,能避开大部分跨设备报错。inductor后端不稳我也遇到过,特别是第二次跑挂的情况,多半是graph缓存和cudnn benchmark冲突,可以试试torch._dynamo.config.suppress_errors=True先跑通再排查。另外如果你用的是HF的模型,可以看看官方有没有专门的compile兼容分支,像llama的某些版本有特定优化补丁。实在不行就退回torchscript或者纯CUDA graph,虽然老套但胜在稳定。
动态shape确实是compile的硬伤,建议先固定到最大长度试试,或者用torch._dynamo的dynamic参数调优。
inductor对变长输入支持还不行,第二次挂大概率是缓存没重建,可以试试关掉模式或者加torch.compile的fullgraph参数。
torch.compile对动态shape的支持确实还比较糙,7B模型输入长度一波动就容易触发recompile或者设备断言,我建议你试试把padding后的长度对齐到8的倍数,同时给compile传dynamic=True参数,能缓解不少。另外inductor后端那个第一次过第二次挂的问题,多半是缓存没清或者cudagraphs和你的自定义算子冲突,关掉mode="reduce-overhead"试试。我最近在项目里干脆用static shape+固定长度截断,推理速度反而更稳,虽然灵活度差了点但生产环境省心。你那个报错是不是在attention里出现的?如果是的话检查下mask的device属性,有时候广播会把tensor带到奇怪的地方去。
这题我蹲过,动态shape直接上fullgraph=True基本必炸,先固定成8的倍数试下吧。
inductor第一次能跑第二次挂大概率是缓存没清,设TORCHINDUCTOR_CACHE_DIR到临时目录看看。
这问题我上周刚踩过一模一样的坑,动态shape在compile下确实容易触发跨设备检查,尤其padding后某些op的tensor shape推断会乱。我目前是把输入长度分桶处理,比如128/256/512这样固定几个档位,配合torch._dynamo的dynamic=True参数,至少能稳定跑起来。inductor后端第一次能过第二次挂我也遇到过,感觉是缓存和CUDA graph的兼容性问题,可以试试加torch.cuda.synchronize或者清一下compiled cache。另外检查下有没有用到torch.where这类容易产生隐式广播的op,改成显式mask会好很多。
动态shape建议用mark_dynamic标记,或者干脆把padding到固定长度再塞给compile,不然跨设备报错真的无解。
inductor二次挂大概率是缓存问题,试试torch._dynamo.config.suppress_errors=True,能苟过就跑起来。
这问题太真实了,动态shape在compile下确实是老大难。我试过把padding到固定长度(比如最大长度的倍数),然后配合torch._dynamo的dynamic=True参数,报错能少很多,但第一次跑还是会慢一点。inductor抽风我也遇过,后来换成cudagraphs配合静态shape才稳下来,不过牺牲了点灵活性。你试试把attention mask也绑进compile的输入里,有时候能绕过那个设备检查的bug。
动态shape这个坑我太熟了,7B模型用compile基本都会撞上。你那个跨设备报错大概率是padding mask没跟着一起走,torch.compile在trace的时候会把某些张量当成静态的,长度一变就崩。我建议你先试试把max_length固定到某个上限,比如2048或者4096,然后padding到那个值,虽然浪费点显存但至少能跑通。另外inductor后端确实有缓存问题,第一次编译完了第二次加载老模型的优化cache可能就失效了,你可以试试设TORCHINDUCTOR_FORCE_DISABLE_CACHE=1看看能不能稳定点。还有个思路是干脆别pad,用可变长度batch然后配合flash attention,但这样compile的支持度更差,我试过几次直接报别的错。说真的,现阶段大模型推理想稳定用compile,要么牺牲灵活性固定shape,要么就老老实实等官方把动态shape支持补全,毕竟这玩意儿还在快速迭代期。你如果主要图推理速度,不如先看看vLLM或者TensorRT-LLM这些专门优化过的框架,自己折腾compile性价比真不高。
这问题我也踩过,动态shape直接静态化或者用mark_dynamic标记下能省不少事,inductor不稳就换cudagraphs试试。
动态shape确实坑,建议把padding长度固定到几个档位,顺便开fullgraph=True,稳定性会好很多。
这问题我也踩过,动态shape在compile下确实容易触发设备断言,尤其padding后某些算子内部会隐式广播。我现在的做法是给模型套一层固定长度包装器,短输入也pad到最大长度,虽然浪费点显存但至少稳定。另外inductor后端建议配一下mode=reduce-overhead,然后确保输入张量先contiguous(),能少很多玄学报错。你试试把dynamic=True参数传给torch.compile,官方说支持动态shape但实际对attention这类算子支持还是半成品。
说实话你这个报错我太熟了,前阵子调一个13B模型也卡在这,最后发现是attention mask和position id在padding后没同步到device上,你检查下是不是某个自定义模块里把tensor显式调了cuda()或者隐式用了CPU上的缓存。动态shape这块torch.compile确实还不太行,它默认会按第一次看到的shape做特化,第二次遇到不同长度就直接崩或者回退,所以要么你固定到最大长度然后mask掉,要么就干脆别用compile,直接用torch.inference_mode加半精度,7B模型在A100上速度也不差多少。
inductor那个“第一次能过第二次挂”我也遇到过,感觉像是guard命中后有些算子被重排,然后某个中间tensor的stride变了,跟后续的view操作冲突。你可以试试加torch._dynamo.config.suppress_errors=True看看能不能退化成eager模式,至少保个底。另外重点查一下模型里有没有用list comprehension或者字典推导式构造tensor的代码,dynamo对这些的追踪经常出幺蛾子,能改成显式torch.cat或者预分配buffer就改一下。
至于最佳实践,我现在的做法是:训练和推理都强制固定seq_len,比如512的倍数,超了直接截断,短了就左边padding,这样compile的稳定性会好很多。如果你实在要变长输入,那就用torch.compile的dynamic=True参数,但别指望它全自动,关键层比如embedding和lm_head还是得手动保证shape一致。还有个小技巧,把模型里面所有.to(x.device)都删掉,改成在forward入口统一处理,能少一半这种跨设备报错。
动态shape确实compile的坑,建议先固定到最大长度加mask试试,inductor对变长支持还不成熟。
之前也踩过这坑,把输入pad到固定长度后稳多了,第一次能过第二次挂多半是缓存问题,清下torch cache再试。
这问题太真实了,动态shape在torch.compile下确实容易触发重编译甚至设备断言,我上次也卡在这。目前比较稳的做法是把输入长度分桶,比如按64或128的倍数padding,然后配合torch._dynamo.config.suppress_errors=True先把报错兜住,至少能跑起来看效果。至于inductor第二次挂,建议试试设置torch._inductor.config.triton.cudagraphs=False,或者干脆用reduce_overhead模式,虽然慢点但稳定很多。另外7B这种规模,如果显存不是特别宽裕,可以考虑把compile范围缩小到decode阶段,prefill保持eager,这样踩坑面积小不少。
说实话你这不算菜,torch.compile对动态shape的支持确实还没到开箱即用的程度,我试过几次也是被跨设备报错折磨得不行。目前最省事的方案就是固定输入长度,或者把padding到某个上限比如2048,然后用static shape模式,这样至少inductor不会在capture时抽风。另外你提到第二次挂的问题,我怀疑是cudagraph缓存或者内存复用导致的,试试torch._dynamo.config.cache_size_limit调大点,或者干脆用mode="reduce-overhead"看看能不能稳一点。要是还不行,可以考虑绕开compile,用vLLM或者TensorRT-LLM这类专门优化推理的框架,省心得多。