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

Detectron2模型转ONNX含pred_masks时报错,求解决方法

Detectron2 Mask RCNN转ONNX含pred_masks报错解决方法

问题场景

需要将Detectron2的Mask RCNN模型转为ONNX格式,要求输出pred_boxes、pred_classes、scores、pred_masks四个字段。当不输出pred_masks时转换完全正常,但包含该字段时会触发ONNX形状推理错误。

预期输出示例:

pred_boxes: Boxes(tensor([[544.5436, 134.7947, 579.4568, 165.4393]], grad_fn=)),
scores: tensor([0.9487],grad_fn=),
pred_classes: tensor([1]),
pred_masks: tensor([[[False, False, False, ..., False, False, False],[False, False, False, ..., False, False, False],[False, False, False, ..., False, False, False]]])

报错信息

UserWarning: The exported ONNX model failed ONNX shape inference. The model will not be executable by the ONNX Runtime. If this is unintended and you believe there is a bug, please report an issue at https://github.com/pytorch/pytorch/issues. Error reported by strict ONNX shape inference: [ShapeInferenceError] Shape inference error(s): (op_type:If, node name: If_2619): [TypeInferenceError] Mismatched tensor element type: source=9 target=2(Triggered internally at ..\torch\csrc\jit\serialization\export.cpp:1421.)_C._check_onnx_proto(proto)

原始问题代码

import cv2
import torch
import numpy as np
from detectron2.config import get_cfg
from detectron2.modeling import build_model
from detectron2.checkpoint import DetectionCheckpointer
from detectron2.data import MetadataCatalog
from detectron2.utils.logger import setup_logger
from detectron2 import model_zoo

cfg = get_cfg()
cfg.merge_from_file("mask_rcnn_X_101_32x8d_FPN_3x.yaml")
cfg.MODEL.WEIGHTS = "model5111_bboxAPs_81_77.pt"
cfg.MODEL.DEVICE = "cpu"
cfg.INPUT.FORMAT = "RGB"
cfg.MODEL.ROI_HEADS.SCORE_THRESH_TEST = (0.6)
cfg.MODEL.NUM_CLASSES = 2
cfg.MODEL.ROI_HEADS.NUM_CLASSES = 2

model = build_model(cfg)
model.eval()
DetectionCheckpointer(model).load(cfg.MODEL.WEIGHTS)

class ModelWrapper(torch.nn.Module):
    def init(self, model):
        super().init()
        self.model = model
    def forward(self, images):
        instances = self.model(images)
        instances = instances[0]
        pred_boxes = instances["instances"].pred_boxes.tensor
        pred_classes = instances["instances"].pred_classes
        scores = instances["instances"].scores
        pred_masks = instances["instances"].pred_masks
        
        print(pred_boxes)
        print(pred_classes)
        print(scores)
        print(pred_masks)
        return (pred_boxes, scores, pred_classes, pred_masks)

model_wrapper = ModelWrapper(model)

inputs = "demo1.png"
im = cv2.imread(inputs)
im = cv2.cvtColor(im, cv2.COLOR_BGR2RGB)
numpy_image = torch.from_numpy(im.transpose(2, 0, 1))
dummy_input = [{"image": numpy_image}]

torch.onnx.export(
    model_wrapper,
    (dummy_input,),
    "converted_model.onnx",
    opset_version=16,
    export_params=True,
    do_constant_folding=False,
    input_names=["input"],
    output_names=["result"],
    dynamic_axes={"input": {0: "batch_size", 2: "height", 3: "width"}, "boxes": {0: "num_boxes"}, "scores": {0: "num_boxes"}, "classes": {0: "num_boxes"}, "masks": {0: "num_boxes", 1: "height", 2: "width"}}
)

print("Model has been converted to ONNX and saved as converted_model.onnx")

修正方案及代码

关键修正点

  1. 修复类初始化方法:将ModelWrapper中的def init(self, model)改为def __init__(self, model),否则类无法正确初始化,导致后续逻辑异常。
  2. 匹配输出名称与动态轴:torch.onnx.export的output_names要对应返回的四个张量,原始代码仅设置["result"],导致动态轴配置的名称无法匹配,触发形状推理错误。
  3. 统一输入张量类型:将numpy转换后的张量转为float类型(Detectron2模型默认输入为float32),避免类型不匹配导致If节点的元素类型错误。
  4. 区分mask维度命名:将动态轴中mask的height/width改为mask_height、mask_width,避免和输入的维度名称混淆。

修正后完整代码

import cv2
import torch
import numpy as np
from detectron2.config import get_cfg
from detectron2.modeling import build_model
from detectron2.checkpoint import DetectionCheckpointer
from detectron2.data import MetadataCatalog
from detectron2.utils.logger import setup_logger
from detectron2 import model_zoo

cfg = get_cfg()
cfg.merge_from_file("mask_rcnn_X_101_32x8d_FPN_3x.yaml")
cfg.MODEL.WEIGHTS = "model5111_bboxAPs_81_77.pt"
cfg.MODEL.DEVICE = "cpu"
cfg.INPUT.FORMAT = "RGB"
cfg.MODEL.ROI_HEADS.SCORE_THRESH_TEST = 0.6
cfg.MODEL.NUM_CLASSES = 2
cfg.MODEL.ROI_HEADS.NUM_CLASSES = 2

model = build_model(cfg)
model.eval()
DetectionCheckpointer(model).load(cfg.MODEL.WEIGHTS)

class ModelWrapper(torch.nn.Module):
    def __init__(self, model):
        super().__init__()
        self.model = model
    def forward(self, images):
        outputs = self.model(images)
        instances = outputs[0]["instances"]
        pred_boxes = instances.pred_boxes.tensor
        pred_classes = instances.pred_classes
        scores = instances.scores
        pred_masks = instances.pred_masks
        
        return pred_boxes, scores, pred_classes, pred_masks

model_wrapper = ModelWrapper(model)

# 预处理输入图像
inputs = "demo1.png"
im = cv2.imread(inputs)
im = cv2.cvtColor(im, cv2.COLOR_BGR2RGB)
numpy_image = torch.from_numpy(im.transpose(2, 0, 1)).float()
dummy_input = [{"image": numpy_image}]

# 导出ONNX
torch.onnx.export(
    model_wrapper,
    (dummy_input,),
    "converted_model.onnx",
    opset_version=16,
    export_params=True,
    do_constant_folding=True,
    input_names=["input"],
    output_names=["boxes", "scores", "classes", "masks"],
    dynamic_axes={
        "input": {0: "batch_size", 2: "height", 3: "width"},
        "boxes": {0: "num_boxes"},
        "scores": {0: "num_boxes"},
        "classes": {0: "num_boxes"},
        "masks": {0: "num_boxes", 1: "mask_height", 2: "mask_width"}
    },
    strict=False
)

print("Model has been converted to ONNX and saved as converted_model.onnx")

验证转换结果

可以用ONNX Runtime加载模型验证:

import onnxruntime as ort
import numpy as np

# 加载模型
sess = ort.InferenceSession("converted_model.onnx")
input_name = sess.get_inputs()[0].name
output_names = [out.name for out in sess.get_outputs()]

# 构造测试输入(通道在前,float32类型)
test_input = np.random.randn(3, 600, 800).astype(np.float32)
outputs = sess.run(output_names, {input_name: test_input})

print("输出boxes形状:", outputs[0].shape)
print("输出scores形状:", outputs[1].shape)
print("输出classes形状:", outputs[2].shape)
print("输出masks形状:", outputs[3].shape)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 17:34:54