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

自定义PyTorch Mask R-CNN导出ONNX:添加batch_size输出维度

问题:为Mask R-CNN的ONNX输出添加batch_size维度

我正尝试将预训练Mask R-CNN模型导出为ONNX格式。该模型基础配置结构已设置batch_size为动态轴,我希望自定义模型,为每个输出添加batch_size维度(即新增一个维度)。我编写了如下代码:

class MaskRCNNModel(torch.nn.Module):
  def __init__(self):
    super(MaskRCNNModel, self).__init__()
    self.model = torchvision.models.detection.maskrcnn_resnet50_fpn(weights='DEFAULT')
    in_features = self.model.roi_heads.box_predictor.cls_score.in_features
    self.model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes=7)
    self.model.load_state_dict(torch.load("saved_dict.torch"))

  def forward(self, input):
    outputs = self.model.forward(input)
    boxes = []
    labels = []
    scores = []
    masks = []
    for result in outputs:
        box, label, score, mask = result.values()
        boxes.append(box)
        labels.append(label)
        scores.append(score)
        masks.append(mask)
    
    return boxes, labels, scores, masks

maskrcnn_model = MaskRCNNModel()
maskrcnn_model.eval()
maskrcnn_model.to(device)

x = torch.rand(1, 3, 512, 512)
x = x.to(device)

maskrcnn_model(x)

torch.onnx.export(maskrcnn_model,
                  x,
                  "base_model_100_epochs.onnx",
                  opset_version=11,
                  input_names=["input"],
                  output_names=["boxes", "labels", "scores", "masks"])

但上述代码未改变导出参数,输出结构仍保持原样。请问应如何修改forward方法,使ONNX模型输出添加batch_size维度?


解决方案

原代码的问题在于,forward方法仅将每个样本的输出收集到列表中,但列表无法被ONNX识别为带batch维度的张量。需要用torch.stack()将列表中的张量拼接,显式新增batch维度。

基础修改(样本输出形状一致时)

修改后的forward方法如下:

def forward(self, input):
    outputs = self.model(input)  # 直接调用model(input)即可,无需显式调用forward
    boxes = []
    labels = []
    scores = []
    masks = []
    for result in outputs:
        box, label, score, mask = result.values()
        boxes.append(box)
        labels.append(label)
        scores.append(score)
        masks.append(mask)
    
    # 使用torch.stack为每个输出添加batch维度,dim=0指定新增维度为第一个维度
    boxes = torch.stack(boxes, dim=0)
    labels = torch.stack(labels, dim=0)
    scores = torch.stack(scores, dim=0)
    masks = torch.stack(masks, dim=0)
    
    return boxes, labels, scores, masks

兼容输出形状不一致的情况(检测框数量不同)

如果batch中每个样本的检测框数量不一致,torch.stack()会报错,需要先对每个样本的输出做padding,确保形状统一后再拼接:

def forward(self, input):
    outputs = self.model(input)
    boxes = []
    labels = []
    scores = []
    masks = []
    
    # 找到batch中最大的检测框数量,作为padding后的统一长度
    max_num_boxes = max(len(result["boxes"]) for result in outputs)
    
    for result in outputs:
        box = result["boxes"]
        label = result["labels"]
        score = result["scores"]
        mask = result["masks"]
        
        pad_num = max_num_boxes - len(box)
        
        # 对各输出做padding,填充值按需设置(比如boxes用0,labels用背景类,scores用0)
        padded_box = torch.nn.functional.pad(box, (0,0,0,pad_num), value=0.0)
        padded_label = torch.nn.functional.pad(label, (0,pad_num), value=0)  # 假设0为背景类
        padded_score = torch.nn.functional.pad(score, (0,pad_num), value=0.0)
        # masks形状为[N,1,H,W],在第0维度padding
        padded_mask = torch.nn.functional.pad(mask, (0,0,0,0,0,0,0,pad_num), value=0.0)
        
        boxes.append(padded_box)
        labels.append(padded_label)
        scores.append(padded_score)
        masks.append(padded_mask)
    
    # 拼接成带batch维度的张量
    boxes = torch.stack(boxes, dim=0)
    labels = torch.stack(labels, dim=0)
    scores = torch.stack(scores, dim=0)
    masks = torch.stack(masks, dim=0)
    
    return boxes, labels, scores, masks

导出ONNX时的动态batch配置

如果需要支持动态batch_size,需在torch.onnx.export中添加dynamic_axes参数:

torch.onnx.export(maskrcnn_model,
                  x,
                  "base_model_100_epochs.onnx",
                  opset_version=11,
                  input_names=["input"],
                  output_names=["boxes", "labels", "scores", "masks"],
                  dynamic_axes={
                      "input": {0: "batch_size"},
                      "boxes": {0: "batch_size"},
                      "labels": {0: "batch_size"},
                      "scores": {0: "batch_size"},
                      "masks": {0: "batch_size"}
                  })

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 07:10:30