You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.07 07:15:34