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

Detectron2 Mask R-CNN转ONNX遇张量类型不匹配错误求解决

Detectron2转ONNX时的If节点类型不匹配问题

转换Detectron2模型为ONNX后运行报错:

Fail: [ONNXRuntimeError] : 1 : FAIL : Load model from /dbfs/redacted/model_simp.onnx failed:Node (If_3504) Op (If) [TypeInferenceError] Mismatched tensor element type: inferred=bool declared=uint8

问题出在模型输出的mask上,推测PyTorch中mask是uint8类型,但ONNX期望bool类型(反之亦然)。去掉mask输出时,ONNX模型能正常运行。

转换代码如下:

import torch
from detectron2.config import get_cfg
from detectron2 import model_zoo
from detectron2.modeling import build_model
from detectron2.checkpoint import DetectionCheckpointer
from detectron2.export import TracingAdapter
import onnx
import cv2
import onnxruntime as ort
import numpy as np

cfg = get_cfg()
model_file_name = "COCO-InstanceSegmentation/mask_rcnn_R_101_FPN_3x.yaml"
cfg.merge_from_file(model_zoo.get_config_file(model_file_name))
cfg.MODEL.DEVICE = 'cpu'
cfg.MODEL.WEIGHTS = 'model_final.pth'
cfg.MODEL.ROI_HEADS.SCORE_THRESH_TEST = 0.5
cfg.MODEL.ROI_HEADS.BATCH_SIZE_PER_IMAGE = 64
cfg.TEST.EVAL_PERIOD = 200
cfg.DATALOADER.NUM_WORKERS = 2
cfg.SOLVER.IMS_PER_BATCH = 2
cfg.INPUT.MASK_FORMAT= 'bitmask'
cfg.SOLVER.BASE_LR = 0.001
cfg.SOLVER.MAX_ITER = 2000
cfg.MODEL.ROI_HEADS.NUM_CLASSES = 3

class TorchModel(torch.nn.Module):
    def __init__(self, cfg) -> None:
        super().__init__()
        self.model = build_model(cfg) # Build Model
        _ = DetectionCheckpointer(self.model).load(cfg.MODEL.WEIGHTS)  # Load weights
        self.model = self.model.eval() # In evaluation mode
    
    def forward(self, INPUT):
        if isinstance(INPUT, (np.ndarray, torch.Tensor)): # it supports just 1 image
            INPUT = [{"image":INPUT}]

        with torch.no_grad():
            outputs = self.model(INPUT)[0]['instances']
        
        boxes, labels, scores, masks = outputs.pred_boxes.tensor, outputs.pred_classes, outputs.scores.detach(), outputs.pred_masks.to(dtype=torch.uint8)

        return boxes, labels, scores, masks
model = TorchModel(cfg)
model = model.eval()

metadata = MetadataCatalog.get(VALID_DATA_SET_NAME)
dataset_valid = DatasetCatalog.get(VALID_DATA_SET_NAME)
im = cv2.imread(dataset_valid[7]['file_name'])
im_torch = torch.as_tensor(im.astype("float32").transpose(2, 0, 1))
inputs = [{"image": im_torch}]

traceable_model = TracingAdapter(model, inputs, None)

torch.onnx.export(traceable_model, (im_torch,),
                  "model.onnx", 
                  opset_version=17,
                  input_names = ['image'],
                  output_names = ['boxes', 'labels', 'scores', 'mask'],
                  dynamic_axes={'image' : {1 : 'height', 2: 'width'},
                                'boxes' : {0 : 'num_boxes'},
                                'labels' : {0 : 'num_boxes'},
                                'scores' : {0 : 'num_boxes'},
                                'masks' : {0 : 'num_boxes', 1 : 'mask_height', 2 : 'mask_width'}
                                })

session = ort.InferenceSession("model.onnx")
im_test = cv2.imread('demo_image.png')
im_test_torch = torch.as_tensor(im_test.astype("float32").transpose(2, 0, 1))

im_test_numpy = im_test_torch.detach().numpy()

outputs = session.run(None, {"image": im_test_numpy})

ONNX模型图的最后几个节点:

Node: Squeeze_3500, OpType: Squeeze, Inputs: ['onnx::Squeeze_3566', 'onnx::Squeeze_3567'], Outputs: ['N']
Node: Constant_3501, OpType: Constant, Inputs: [], Outputs: ['onnx::Equal_3569']
Node: Equal_3502, OpType: Equal, Inputs: ['N', 'onnx::Equal_3569'], Outputs: ['onnx::Cast_3570']
Node: Cast_3503, OpType: Cast, Inputs: ['onnx::Cast_3570'], Outputs: ['onnx::If_3571']
Node: If_3504, OpType: If, Inputs: ['onnx::If_3571'], Outputs: ['bitmasks']
Node: /model/model/Cast_11, OpType: Cast, Inputs: ['bitmasks'], Outputs: ['/model/model/Cast_11_output_0']
Node: /model/Cast, OpType: Cast, Inputs: ['/model/model/Cast_11_output_0'], Outputs: ['mask']
解决方案

1. 调整mask类型转换的时机和方式

在模型forward方法中,提前明确mask的类型,帮助TorchScript追踪时生成正确的ONNX节点:

# 修改forward中的mask提取代码
# 若期望输出bool类型
masks = outputs.pred_masks.bool()
# 若坚持输出uint8类型,添加显式声明
masks = torch.as_tensor(outputs.pred_masks.to(dtype=torch.uint8), dtype=torch.uint8)

2. 强制指定ONNX导出的输出类型

在torch.onnx.export中通过output_types参数明确每个输出的张量类型,让ONNX生成匹配的类型声明:

torch.onnx.export(traceable_model, (im_torch,),
                  "model.onnx", 
                  opset_version=17,
                  input_names = ['image'],
                  output_names = ['boxes', 'labels', 'scores', 'mask'],
                  dynamic_axes={'image' : {1 : 'height', 2: 'width'},
                                'boxes' : {0 : 'num_boxes'},
                                'labels' : {0 : 'num_boxes'},
                                'scores' : {0 : 'num_boxes'},
                                'masks' : {0 : 'num_boxes', 1 : 'mask_height', 2 : 'mask_width'}
                                },
                  output_types=[torch.float32, torch.int64, torch.float32, torch.uint8])

3. 使用Detectron2官方导出工具

官方工具会自动处理Detectron2模型的类型兼容问题,比手动追踪更可靠:

from detectron2.export import export_model_to_onnx

export_model_to_onnx(
    cfg,
    "model_final.pth",
    "model.onnx",
    opset_version=17,
    input_shape=(3, None, None)
)

4. 手动修正已导出的ONNX模型

如果已经生成了模型,可直接修改ONNX文件中If节点的输出类型声明:

import onnx

model = onnx.load("model.onnx")
# 找到目标If节点并修改输出类型
for node in model.graph.node:
    if node.name == "If_3504":
        for output_name in node.output:
            for value_info in model.graph.value_info:
                if value_info.name == output_name:
                    # 根据实际需求选择UINT8或BOOL
                    value_info.type.tensor_type.elem_type = onnx.TensorProto.UINT8
onnx.save(model, "model_fixed.onnx")

内容的提问来源于stack exchange,提问作者Yannick

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 14:24:49