PyTorch模型torch.jit.trace时报aten::fill_操作缺失错误如何解决
修复方案
方案一:修改赋值逻辑适配trace
把直接赋值原生bool的操作改为显式传入PyTorch scalar,修改问题行即可:
# 替换x_mask赋值行 x_mask[crop_h[0] : crop_h[1], crop_w[0] : crop_w[1]] = torch.tensor(True, device=x.device, dtype=torch.bool) # 替换s_mask赋值行 s_mask[crop_sh[0] : crop_sh[1], crop_sw[0] : crop_sw[1]] = torch.tensor(True, device=s.device, dtype=torch.bool)
该方案仅修改赋值逻辑,不需要调整其他代码,即可正常完成trace。注意如果后续输入尺寸变化,需要重新trace。
方案二:改用torch.jit.script替代trace
如果你的业务场景存在动态输入尺寸的需求,更推荐使用支持动态逻辑的torch.jit.script,不需要修改原有业务代码,直接替换trace调用即可:
# 把原来的trace行替换为script scripted_cell = torch.jit.script(patt) print(scripted_cell)
script会直接解析Python代码语法生成静态图,天然支持条件分支、动态形状逻辑,兼容性更好。
内容的提问来源于stack exchange,提问作者mazy
相关产品推荐
相关产品推荐

