PyTorch模型转TorchScript时Runtime Error问题求解
PyTorch模型转TorchScript时Concat类报错解决方案
问题代码
import torch import torch.nn as nn class Concat(nn.Module): def __init__(self): super(Concat, self).__init__() def forward(self, x): return torch.cat(x,1) class Net(nn.Module): def __init__(self) -> None: super().__init__() self.conv1 = nn.Conv2d(3, 16, 3, 1) self.conv2 = nn.Conv2d(16, 32, 3, 1) def forward(self, x): y = self.conv1(x) y = self.conv2(y) z = self.conv1(x) z = self.conv2(z) return (y, z) net = nn.Sequential( Net(), Concat() ) mobile_net = torch.quantization.convert(net) scripted_net = torch.jit.script(mobile_net)
报错信息
RuntimeError Traceback (most recent call last) Cell In [2], line 26 21 net = nn.Sequential( 22 Net(), 23 Concat() 24 ) 25 mobile_net = torch.quantization.convert(net) ---> 26 scripted_net = torch.jit.script(mobile_net) File ~\anaconda3\envs\yolov5pytorch\lib\site-packages\torch\jit\_script.py:1286, in script(obj, optimize, _frames_up, _rcb, example_inputs) 1284 if isinstance(obj, torch.nn.Module): 1285 obj = call_prepare_scriptable_func(obj) -> 1286 return torch.jit._recursive.create_script_module( 1287 obj, torch.jit._recursive.infer_methods_to_compile 1288 ) 1290 if isinstance(obj, dict): 1291 return create_script_dict(obj) File ~\anaconda3\envs\yolov5pytorch\lib\site-packages\torch\jit\_recursive.py:476, in create_script_module(nn_module, stubs_fn, share_types, is_tracing) 474 if not is_tracing: 475 AttributeTypeIsSupportedChecker().check(nn_module) --> 476 return create_script_module_impl(nn_module, concrete_type, stubs_fn) File ~\anaconda3\envs\yolov5pytorch\lib\site-packages\torch\jit\_recursive.py:538, in create_script_module_impl(nn_module, concrete_type, stubs_fn) 535 script_module._concrete_type = concrete_type 537 # Actually create the ScriptModule, initializing it with the function we just defined --> 538 script_module = torch.jit.RecursiveScriptModule._construct(cpp_module, init_fn) 540 # Compile methods if necessary 541 if concrete_type not in concrete_type_store.methods_compiled: File ~\anaconda3\envs\yolov5pytorch\lib\site-packages\torch\jit\_script.py:615, in RecursiveScriptModule._construct(cpp_module, init_fn) 602 """ 603 Construct a RecursiveScriptModule that's ready for use. PyTorch 604 code should use this to construct a RecursiveScriptModule instead (...) 612 init_fn: Lambda that initializes the RecursiveScriptModule passed to it. 613 """ 614 script_module = RecursiveScriptModule(cpp_module) -> 615 init_fn(script_module) 617 # Finalize the ScriptModule: replace the nn.Module state with our 618 # custom implementations and flip the _initializing bit. 619 RecursiveScriptModule._finalize_scriptmodule(script_module) File ~\anaconda3\envs\yolov5pytorch\lib\site-packages\torch\jit\_recursive.py:516, in create_script_module_impl.<locals>.init_fn(script_module) 513 scripted = orig_value 514 else: 515 # always reuse the provided stubs_fn to infer the methods to compile --> 516 scripted = create_script_module_impl(orig_value, sub_concrete_type, stubs_fn) 518 cpp_module.setattr(name, scripted) 519 script_module._modules[name] = scripted File ~\anaconda3\envs\yolov5pytorch\lib\site-packages\torch\jit\_recursive.py:542, in create_script_module_impl(nn_module, concrete_type, stubs_fn) 540 # Compile methods if necessary 541 if concrete_type not in concrete_type_store.methods_compiled: --> 542 create_methods_and_properties_from_stubs(concrete_type, method_stubs, property_stubs) 543 # Create hooks after methods to ensure no name collisions between hooks and methods. 544 # If done before, hooks can overshadow methods that aren't exported. 545 create_hooks_from_stubs(concrete_type, hook_stubs, pre_hook_stubs) File ~\anaconda3\envs\yolov5pytorch\lib\site-packages\torch\jit\_recursive.py:393, in create_methods_and_properties_from_stubs(concrete_type, method_stubs, property_stubs) 390 property_defs = [p.def_ for p in property_stubs] 391 property_rcbs = [p.resolution_callback for p in property_stubs] --> 393 concrete_type._create_methods_and_properties(property_defs, property_rcbs, method_defs, method_rcbs, method_defaults) RuntimeError: Arguments for call are not valid. The following variants are available: aten::cat(Tensor[] tensors, int dim=0) -> Tensor: Expected a value of type 'List[Tensor]' for argument 'tensors' but instead found type 'Tensor (inferred)'. Inferred the value for argument 'tensors' to be of type 'Tensor' because it was not annotated with an explicit type. aten::cat.names(Tensor[] tensors, str dim) -> Tensor: Expected a value of type 'List[Tensor]' for argument 'tensors' but instead found type 'Tensor (inferred)'. Inferred the value for argument 'tensors' to be of type 'Tensor' because it was not annotated with an explicit type. aten::cat.names_out(Tensor[] tensors, str dim, *, Tensor(a!) out) -> Tensor(a!): Expected a value of type 'List[Tensor]' for argument 'tensors' but instead found type 'Tensor (inferred)'. Inferred the value for argument 'tensors' to be of type 'Tensor' because it was not annotated with an explicit type. aten::cat.out(Tensor[] tensors, int dim=0, *, Tensor(a!) out) -> Tensor(a!): Expected a value of type 'List[Tensor]' for argument 'tensors' but instead found type 'Tensor (inferred)'. Inferred the value for argument 'tensors' to be of type 'Tensor' because it was not annotated with an explicit type. The original call is: File "C:\Users\pawan\AppData\Local\Temp\ipykernel_16484\3929675973.py", line 6 def forward(self, x): return torch.cat(x,1) ~~~~~~~~~ <--- HERE
问题原因
TorchScript的静态类型推断无法自动识别Concat类forward方法的输入x是Tensor元组/列表,错误将其推断为单个Tensor类型,导致torch.cat接收到的参数类型不符合要求(torch.cat需要传入Tensor的序列,而非单个Tensor)。
而将torch.cat整合到Net类的forward方法中能正常运行,是因为在Net内部,TorchScript可以明确追踪到返回的(y,z)是Tensor元组,能正确推断torch.cat的输入类型。
解决方案
方案1:添加类型注解
给Concat的forward方法显式标注输入类型,帮助TorchScript正确推断:
import torch import torch.nn as nn from typing import Tuple class Concat(nn.Module): def __init__(self): super(Concat, self).__init__() def forward(self, x: Tuple[torch.Tensor, torch.Tensor]) -> torch.Tensor: return torch.cat(x, 1)
如果需要支持任意数量的Tensor输入,可以用Tuple[torch.Tensor, ...]或List[torch.Tensor]作为类型注解。
方案2:显式转换输入为序列
在forward方法中手动将输入转为列表,强制TorchScript识别输入为Tensor序列:
class Concat(nn.Module): def __init__(self): super(Concat, self).__init__() def forward(self, x): # 显式将输入转为列表,明确告诉TorchScript这是Tensor序列 return torch.cat(list(x), 1)
方案3:使用TorchScript类型提示
如果不想引入typing模块,可以用torch.jit.annotate标注类型:
class Concat(nn.Module): def __init__(self): super(Concat, self).__init__() def forward(self, x): x = torch.jit.annotate(tuple[torch.Tensor, torch.Tensor], x) return torch.cat(x, 1)
修改后重新运行代码,即可成功将模型转换为TorchScript。
内容的提问来源于stack exchange,提问作者Pawan Kumar J
相关产品推荐
相关产品推荐

