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
相关产品推荐
相关产品推荐

