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

目标检测训练遇RuntimeError:张量尺寸不匹配求解决

目标检测自定义数据集训练报错解决

问题场景

自定义了ReceiptDataset用于收据目标检测任务,实例化数据集并启动训练循环时,出现张量尺寸不匹配的RuntimeError,报错信息显示不同样本的目标框张量尺寸不一致(如[11,4]和[9,4])。

自定义ReceiptDataset代码

from torch.nn.utils.rnn import pad_sequence
import torch.nn.functional as F
import os
import cv2
import numpy as np
import torch

class ReceiptDataset(torch.utils.data.Dataset):
    def __init__(self, train_dir, width, height, labels, transforms=None):
        self.images = os.listdir(train_dir)
        self.width = width
        self.height = height
        self.train_dir = train_dir
        self.labels = labels
        self.transforms = transforms

    def __getitem__(self, idx):
        img_name = self.images[idx]
        img_path = os.path.join(self.train_dir, img_name)

        img = cv2.imread(img_path)
        img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB).astype(np.float32)
        img_res = cv2.resize(img_rgb, (self.width, self.height), cv2.INTER_AREA)

        img_res /= 255.0

        annot = self.labels[str(img_name)]

        lbls = []
        boxes = []
        target = {}

        ht, wt, _ = img.shape
        
        for item in annot:
            x, y, box_wt, box_ht, lbl = item

            x_min = x
            x_max = x + box_wt
            y_min = y
            y_max = y + box_ht

            x_min_corr = (x_min / wt) * self.width
            x_max_corr = (x_max / wt) * self.width
            y_min_corr = (y_min / ht) * self.height
            y_max_corr = (y_max / ht) * self.height

            boxes.append([x_min_corr, y_min_corr, x_max_corr, y_max_corr])
            lbls.append(classes.index(str(lbl)))

        boxes = torch.as_tensor(boxes, dtype=torch.float32)
        lbls = torch.as_tensor(lbls, dtype=torch.int64)

        area = (boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0])
        iscrowd = torch.zeros((boxes.shape[0],), dtype=torch.int64)

        target["boxes"] = boxes
        target["labels"] = lbls
        target["image_id"] = torch.as_tensor(idx)
        target["area"] = area
        target["iscrowd"] = iscrowd

        if self.transforms:
            trans = self.transforms(image=img_res,
                                    bboxes=target["boxes"],
                                    labels=lbls)
            img_res = trans["image"]
            target["boxes"] = torch.Tensor(trans["bboxes"])

        return img_res, target

    def __len__(self):
        return len(self.images)

数据集实例化代码

train_dataset = ReceiptDataset("label-detector/images", width, height, plabels)

训练代码片段

from engine import train_one_epoch, evaluate

for epoch in range(num_epochs):
    train_one_epoch(model, optim, train_loader, device, epoch, print_freq=2)
    lr_scheduler.step()
    evaluate(model, test_loader, device)

报错信息

RuntimeError: stack expects each tensor to be equal size, but got [11,4] at entry 0 and [9,4] at entry 1

问题原因及解决办法

原因

PyTorch默认的DataLoader会尝试将batch内的所有张量堆叠(stack)成统一尺寸的张量,但目标检测任务中每张图片的目标框、标签数量不一致,导致堆叠失败。

解决步骤

  1. 自定义collate_fn函数:
    该函数负责处理batch内的数据,不对不同长度的目标张量强行堆叠,同时修正图像的维度顺序(PyTorch模型期望输入为(C,H,W),而OpenCV读取的是(H,W,C))。

    def collate_fn(batch):
        images = []
        targets = []
        for img, target in batch:
            # 转换图像维度从(H,W,C)到(C,H,W)
            images.append(img.permute(2, 0, 1))
            targets.append(target)
        # 图像尺寸统一,可以直接堆叠
        images = torch.stack(images, dim=0)
        return images, targets
    
  2. DataLoader中指定collate_fn:
    在创建训练和测试数据加载器时,传入自定义的collate_fn:

    train_loader = torch.utils.data.DataLoader(
        train_dataset,
        batch_size=8,  # 根据硬件调整batch size
        shuffle=True,
        collate_fn=collate_fn
    )
    
    test_loader = torch.utils.data.DataLoader(
        test_dataset,
        batch_size=8,
        shuffle=False,
        collate_fn=collate_fn
    )
    

额外注意点

  • 不要尝试用lbls += [-1] * (NUM_CLASSES - len(lbls))固定标签数量,这不符合目标检测的标注逻辑(每个标签对应一个目标框,而非每个类别对应一个标签)。
  • 若使用albumentations等数据增强库,确保transforms返回的bboxes格式与模型期望一致(如[xmin, ymin, xmax, ymax])。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 08:25:58