最近在试着用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下基本就是灾难,你那个跨设备报错八成是padding后某些op在graph里被错误地按固定shape特化了。我现在做法是分桶处理,把输入长度按64的倍数分成几个档位,每个桶单独compile,虽然显存占用多点但稳定多了。inductor后端那个第二次挂的问题我也遇到过,感觉是cudagraph的缓存和动态shape冲突了,试试在compile时加mode="reduce-overhead"再配个torch._dynamo.config.suppress_errors=True,至少能让它回退到eager模式。另外你如果用的是llama结构,试试把attention里的mask改成固定最大长度的预计算版本,别在forward里动态生成,能省不少麻烦。还有个小技巧,实在搞不定动态shape就把输入pad到模型训练时的max_seq_len,用torch.compile的fullgraph=True强制静态化,推理速度虽然不一定最快但至少不报错。最后想问下你用的是哪个版本的torch,2.1之后对动态shape的支持好了不少,如果还在用2.0.1建议先升级试试。
这问题我上周刚踩过,动态shape在compile下确实容易触发device断言,尤其是padding后attention mask参与计算的时候。我的做法是干脆把输入长度分桶,比如64、128、256这样,每个桶固定max_len再compile,虽然牺牲点显存但稳很多。inductor第二次挂的话试试加mode=“max-autotune”然后关掉dynamic=True,或者干脆回退到reduce-overhead,别追求极致编译。另外你检查下模型里有没有用到tensor.size()直接做切片的操作,改成mask或固定shape能绕开不少坑。
说实话你这个坑我上个月也踩过,7B模型动态shape下compile基本就是地狱模式。我后来是直接把padding到固定长度,比如512或1024的倍数,然后配合torch._dynamo.mark_dynamic标注关键维度,报错能少一大半。但你说的第二次推理挂掉我也遇到过,感觉inductor的缓存和CUDA graph的replay在某些算子组合下就是有bug,尤其当模型里有条件分支或者循环依赖tensor shape时,建议把mode设成reduce-overhead或者干脆用max-autotune-no-cudagraphs试试。
另外还有个思路,如果你输入长度变化真的很大,不如别硬刚compile,用vLLM或者FasterTransformer那个思路,把attention和FFN拆开手动优化,或者直接上TensorRT-LLM,虽然配置麻烦点但稳定很多。我试过在compile环境下把动态维度改成静态,然后通过多次不同长度的profile来预热,效果比动态编译好,但启动时间就上去了。
还有个小坑,checkpoint加载的时候如果模型里有meta device之类的操作,compile会更容易跨设备报错,试试先load到CPU再转GPU。你现在用的CUDA版本和torch版本是多少?我这边2.0.1+cu118配dynamo有时候也会抽风,后来升到2.1.0就顺了一丢丢,但也没完全根治。说实话这功能还是太新,生产环境真不太敢依赖。
这问题太真实了,我前几天在llama.cpp和torch.compile之间反复横跳也差点砸键盘。动态shape确实是compile的坑,我后来是把pad到固定长度加上max_length限制才稳定下来,虽然牺牲点吞吐但至少不炸了。inductor那个玄学问题我也遇到过,怀疑是缓存了某些图的优化结果,第二次输入shape变了就冲突,建议试试dynamic=True参数或者干脆每次推理前清一下torch的cache。另外你如果主要是要加速生成,其实可以考虑用vLLM或者TGI,纯torch compile在大模型场景里投入产出比真的一般。
这问题我熟,7B模型跑compile大概率就是栽在动态shape上,torch.compile对变化维度特别敏感。你试试把输入长度固定到最大,或者用torch._dynamo.mark_dynamic把关键维度标一下,能省不少事。inductor不稳定确实常见,尤其是带cache的二次推理,可以试试关掉cudagraphs或者用max-autotune模式,虽然慢点但稳。还有个小技巧,把KV cache的tensor预先分配好,别让runtime去动态扩展,跨设备报错能少一大半。
这问题我也踩过,动态shape在compile下确实容易触发跨设备检查,尤其padding后attention mask的维度变化。你试试把输入长度分桶(比如256、512、1024),然后用torch._dynamo.mark_dynamic或者干脆固定到几个档位,能避开不少雷。inductor第一次过第二次挂多半是缓存或graph重编译的问题,建议关掉cudagraphs或者设torch.compile(mode="max-autotune-no-cudagraphs")看看,我这边这么搞稳定很多。另外7B模型建议配合flash-attention的varlen接口,比padding干净得多。
动态shape确实容易踩坑,先pad到固定长度试试,另外inductor对变长输入兼容性还不太行。
遇到过同样问题,动态shape直接上compile就是折磨,建议先固定长度跑通再优化。
试试把padding到固定长度关掉cudnn benchmark,或者用torch._dynamo的dynamic参数,能省不少事。
这问题我太有共鸣了,torch.compile对动态shape的容忍度确实低得让人头秃。你那个跨设备报错我猜是某个子模块被重编译后,缓存里的图跟实际输入对不上,导致张量被分到不同流处理器上了。我试过的办法是给输入长度设个上限,然后统一pad到那个固定值,虽然浪费点显存但至少稳定,比它自己动态recompile省心多了。另外你可以试试把dynamic=True参数显式传给torch.compile,让它提前知道shape会变,我用了这个之后报错频率低了很多,但偶尔还是会抽风。至于inductor后端第一次能过第二次挂,我怀疑是CUDA graph捕获和某些算子不兼容,你可以试试把mode改成reduce-overhead或者关掉cudagraphs,牺牲点速度换稳定。还有个歪招,干脆用ONNX导出再走TensorRT,虽然折腾但推理起来是真稳,就是调试周期长。说实话现在大模型推理真没到无脑compile就能提速的阶段,我最后是直接退回torch.jit.script,把几个关键瓶颈层手写优化,效果反而比全量compile靠谱。
这问题我踩过一模一样的坑,动态shape在compile下确实容易触发设备检查的边界情况,尤其padding后某些op的tensor shape推断会漂。我后来是把输入长度按桶(比如64的倍数)做静态化,配合torch._dynamo的dynamic=True参数,基本稳定了。另外inductor第一次能过第二次挂,大概率是guard缓存和cudnn benchmark的交互问题,试试torch._inductor.compile(..., mode="max-autotune")或者关掉cudnn benchmark,能缓解不少。不过说实话,7B这种规模,compile收益有限,还不如先优化kv cache和attention实现。
试试把动态维度的最大值固定下来,配合torch._dynamo.mark_dynamic,能少踩不少坑。
动态shape就是compile的痛点,建议先固定到最大长度+mask,或者试试dynamo的dynamic=True参数。
这问题太真实了,动态shape在compile下确实容易触发跨device的bug,尤其7B这种模型里embedding和lm_head的权重共享,padding一变就容易踩到隐式广播的坑。我目前是把输入长度分桶处理,比如256、512、1024这几个档位,每个桶单独compile,既保住性能又避开动态shape的雷。inductor二次挂的问题我也遇到过,后来发现是cudagraph的缓存和可变长度冲突,建议直接关掉mode=“reduce-overhead”试试,或者干脆用torch._dynamo的dynamic=True参数,但得配合mark_dynamic手动指定维度,不然还是容易飘。另外你可以试试把padding放到模型的最后一层再处理,或者用左padding,有时候能绕开一些奇怪的device断言。
这问题我上周刚踩完坑,动态shape确实会让inductor生成跨设备的临时张量,尤其是注意力mask那块。我现在的做法是直接把输入pad到固定长度(比如512的倍数),然后用torch._inductor.config.dynamic_shapes=False关掉动态支持,虽然会浪费点显存但至少稳。另外你试过把max_length设成模型实际能接受的上限,然后配合静态缓存吗?第二次挂的那个问题我怀疑是CUDA graph缓存没刷新,可以用torch._inductor.config.triton.cudagraphs = False先排除下。
这问题我上周刚踩完坑,动态shape确实会让inductor生成跨设备的临时张量,尤其是注意力mask那块。我现在的做法是直接把输入pad到固定长度(比如512的倍数),然后用torch._inductor.config.dynamic_shapes=False关掉动态支持,虽然会浪费点显存但至少稳。另外你试过把max_length设成模型实际能接受的上限,然后配合静态缓存吗?第二次挂的那个问题我怀疑是CUDA graph缓存没刷新,可以用torch._inductor.config.triton.cudagraphs = False先排除下。
说实话你这不算菜,动态shape跟compile本来就是老冤家了,PyTorch官方文档里都写了要尽量用static shape,你试试把padding那步挪到模型外面做,或者用torch.compile的dynamic=True参数,虽然会牺牲一点性能但至少能跑起来。我上次搞一个6B模型也遇到类似问题,后来发现是attention mask的维度没跟着input一起更新,建议你检查一下所有参与矩阵乘法的张量是不是都严格对齐了device和shape。另外inductor后端那个第一次能过第二次挂的问题,我怀疑跟CUDA graph缓存有关,你可以试试设置torch._inductor.config.triton.cudagraphs=False,虽然慢点但稳定很多。还有个野路子,输入长度如果变化范围可控的话,干脆把长度聚成几个档位,每个档位单独编译一个graph,这样既避免动态shape又不用牺牲太多性能。最后想问你用的是哪个版本的transformers,有些老版本跟2.0的compile兼容性特别差,升级到最新版可能直接就好了。
动态shape确实是compile的坑,固定长度加mask能稳不少,但牺牲点吞吐也能忍。
这问题我也踩过,动态shape在compile下确实容易触发跨设备检查,建议先把max_length固定住或者用static cache,能省不少事。inductor第一次过第二次挂八成是cudagraphs和动态shape的兼容性问题,可以试试把mode设成reduce-overhead或者直接关掉cudagraphs看看。另外7B模型跑compile收益其实有限,瓶颈多在显存带宽,不如先优化一下kv cache和attention实现。
动态shape确实是torch.compile目前的老大难,特别是LLM这种变长输入场景。我自己的经验是先把padding到固定长度(比如取训练时的max_len),然后配合torch._dynamo.mark_dynamic(或者干脆用torch.utils.data下的固定batch)能解决大部分跨设备报错。不过你说的inductor二次推理挂掉也遇到过,感觉是capture graph时某些op被错误缓存了,这时候可以试试关掉mode="reduce-overhead"或者加个torch._dynamo.config.suppress_errors=True先让代码跑起来再排查。另外7B模型建议优先用reduce-overhead配合cudagraphs,但前提是输入shape必须完全静态,否则真的容易玄学。还有个偏方:把模型里所有reshape/view操作都换成torch.permute+contiguous,有时候能避开device断言。别灰心,这坑我踩了三周才稳定。
对了,你试过用compile的fullgraph=True吗?有时候能暴露更多graph break的具体位置。之前看到群里有人用torch._dynamo.config.capture_scalar_outputs=True解决过类似问题,但不确定对7B是否有效。如果还不行,可以退而求其次用torch.jit.script冻结某些层,虽然性能提升小但至少不崩。说到底现在compile对大模型推理的收益还没到质变,能跑通就是胜利。
这问题我也踩过,动态shape在compile下确实容易触发跨device的bug,尤其padding后某些算子会隐式访问未初始化的位置。我的做法是先对输入长度做分桶(比如按32的倍数截断),再用static shape跑compile,效果比直接上动态shape稳很多。inductor第一次过第二次挂的玄学我猜是graph缓存和内存池没对齐,试试torch._dynamo.config.suppress_errors=True先跑通再说,虽然会牺牲点性能但至少不崩。另外7B模型其实可以考虑用vLLM或者TGI那套现成方案,自己调compile的投入产出比可能不太划算。
说实话你这不算菜,动态shape和compile本来就不太对付,7B模型能跑起来已经很能折腾了。我建议先试试torch._dynamo.mark_dynamic给关键输入打个标记,或者干脆把padding长度固定到训练时的max,别让shape在中间层乱跳。inductor第一次成功第二次挂我也遇到过,多半是缓存或graph break的问题,可以试着把mode设成reduce-overhead,再不行就回退到默认后端跑几次看稳定性。另外检查下模型里有没有reshape/view这种隐式改shape的操作,有时候跨设备就是这些地方触发的。