保存TorchScript模块时遇RuntimeError: strides()调用未定义张量
问题:Torch JIT Script保存时报错RuntimeError: strides() called on an undefined Tensor
复现代码
model.eval() trace_script_module = torch.jit.script(model, (image1, image2)) trace_script_module.save("a.pt")
报错信息
RuntimeError: strides() called on an undefined Tensor
已做尝试
- 将模型和输入迁移至CPU,问题未解决
- 删除
trace_script_module.save("a.pt")后无报错,确认是save操作触发的问题,但找不到代码中strides()的调用位置
解决思路与方案
1. 排查模型内未初始化的Tensor
这个错误核心原因是模型中存在未被正确初始化的Tensor——JIT在序列化(save)时会遍历模型所有参数、缓存Tensor,一旦碰到未定义的就会触发该报错。
- 检查模型
__init__方法:确保所有声明的Tensor都有明确的初始化值,不要出现只定义变量名但未赋值的情况 - 检查forward函数:有没有条件分支(比如if-else)导致某些Tensor仅在部分分支中被创建,后续却被全局使用
2. 尝试用torch.jit.trace替代script
如果你的模型没有复杂的动态控制流(比如大量基于张量值的if判断、动态循环),可以换用trace方式生成JIT模块:
model.eval() with torch.no_grad(): trace_script_module = torch.jit.trace(model, (image1, image2)) trace_script_module.save("a.pt")
trace是通过追踪模型执行过程生成IR,对未初始化Tensor的敏感度更低,能规避部分script的序列化问题。
3. 手动清理模型缓存Tensor
在save前手动清理可能存在的无效缓存Tensor:
model.eval() trace_script_module = torch.jit.script(model, (image1, image2)) # 遍历模块,删除未定义的参数/缓存 for module in trace_script_module.modules(): for name, attr in list(module.named_attributes(recurse=False)): if isinstance(attr, torch.Tensor) and not attr.is_defined(): delattr(module, name) trace_script_module.save("a.pt")
4. 验证输入Tensor的合法性
确认image1和image2是完全有效的Tensor,没有空值或未定义维度:
print("image1状态:", image1.is_defined(), "shape:", image1.shape) print("image2状态:", image2.is_defined(), "shape:", image2.shape)
如果存在未定义的输入Tensor,先修复输入再尝试序列化。
内容的提问来源于stack exchange,提问作者hellosleepwalker
相关产品推荐
相关产品推荐

