PyTorch DataLoader加载不等数量边界框报错解决方案咨询
问题描述
我正在为包含不等数量车辆边界框的图像数据集构建PyTorch DataLoader,当batch_size设为1时运行正常,但设为2时出现RuntimeError: stack expects each tensor to be equal size错误。
数据集类代码
class AGR_Dataset(Dataset): def __init__(self, annotations_root, img_root, transform=None): """ Arguments: annotations_root (string): Path to the csv file with annotations. img_root (string): Directory with all the images. transform (callable, optional): Optional transform to be applied on a sample. """ self.annotations_root = annotations_root self.img_root = img_root self.transform = transform def __len__(self): return len(self.annotations_root) def __getitem__(self, idx): # idx may be the index or image name, I think image naem if torch.is_tensor(idx): idx = idx.tolist() idx_name = os.listdir(self.img_root)[idx] # print(idx_name) img_name = os.path.join(self.img_root, idx_name) annotation_data = os.path.join(self.annotations_root, f"{idx_name.removesuffix('.jpg')}.txt") # print(img_name, annotation_data) image = io.imread(img_name) with open(annotation_data, 'r') as file: lines = file.readlines() img_data = [] img_labels = [] for line in lines: line = line.split(',') line = [i.strip() for i in line] line = [float(num) for num in line[0].split()] img_labels.append(int(line[0])) img_data.append(line[1:]) boxes = tv_tensors.BoundingBoxes(img_data, format='CXCYWH', canvas_size=(image.shape[0], image.shape[1])) # sample = {'image': image, 'bbox': boxes, 'labels': img_labels} sample = {'image': image, 'bbox': boxes} if self.transform: sample = self.transform(sample) print(sample['image'].shape) print(sample['bbox'].shape) # print(sample['labels'].shape) return sample
Transform及DataLoader配置
data_transform = v2.Compose([ v2.ToImage(), # v2.Resize(680), v2.RandomResizedCrop(size=(680, 680), antialias=True), # v2.ToDtype(torch.float32, scale=True), v2.ToTensor() ]) transformed_dataset = AGR_Dataset(f'{annotations_path}/test/', f'{img_path}/test/', transform=data_transform) dataloader = DataLoader(transformed_dataset, batch_size=2, shuffle=False, num_workers=0)
遍历DataLoader代码
for i, sample in enumerate(dataloader): print(i, sample) print(i, sample['image'].size(), sample['bbox'].size()) if i == 4: break
错误信息
torch.Size([3, 680, 680]) torch.Size([12, 4]) torch.Size([3, 680, 680]) torch.Size([259, 4]) RuntimeError: stack expects each tensor to be equal size, but got [12, 4] at entry 0 and [259, 4] at entry 1
疑问
- 我认为错误源于图像间边界框数量不等,该如何解决?
- 我的Transform中是否需要ToTensor?v2已使用ToImage(),ToTensor似乎已过时。
已尝试方法
- 注释tv_tensors.BoundingBoxes代码,但Resize无法正常工作;
- 将图像和边界框拆分为sample和target,未解决问题。
解决方案
问题1:解决边界框数量不等导致的batch堆叠错误
PyTorch默认DataLoader会尝试将batch内所有tensor堆叠成统一形状,但不同图像的边界框数量不同,无法直接堆叠。解决核心是自定义collate函数,灵活处理不同长度的边界框数据:
方法1:自定义collate_fn(推荐,适配目标检测场景)
修改DataLoader初始化,传入自定义collate函数,将边界框保留为列表形式:
def collate_fn(batch): # 分离batch中的图像和边界框 images = [item['image'] for item in batch] bboxes = [item['bbox'] for item in batch] # 图像形状统一,可直接堆叠成tensor images = torch.stack(images, dim=0) # 返回字典,边界框以列表形式保存(每个元素对应单张图的边界框) return {'image': images, 'bbox': bboxes} # 初始化DataLoader时传入自定义collate_fn dataloader = DataLoader(transformed_dataset, batch_size=2, shuffle=False, num_workers=0, collate_fn=collate_fn)
遍历DataLoader时,sample['bbox']会是一个列表,每个元素对应单张图像的边界框tensor,不会触发堆叠错误。
方法2:边界框padding(适配需要固定输入形状的模型)
如果模型要求输入固定形状的边界框tensor,可以对每个样本的边界框进行padding,补足到当前batch的最大数量:
def collate_fn_pad(batch): images = [item['image'] for item in batch] bboxes = [item['bbox'] for item in batch] # 找到当前batch中边界框的最大数量 max_num_boxes = max(len(box) for box in bboxes) # 对边界框进行padding,填充值设为0(或其他无效标记) padded_bboxes = [] for box in bboxes: pad_num = max_num_boxes - len(box) padded = torch.nn.functional.pad(box, (0,0,0,pad_num), value=0) padded_bboxes.append(padded) images = torch.stack(images, dim=0) padded_bboxes = torch.stack(padded_bboxes, dim=0) return {'image': images, 'bbox': padded_bboxes}
注意:使用此方法时,模型需要能区分真实边界框和padding的无效框(比如通过标签或额外的mask)。
问题2:Transform中ToTensor的必要性
在TorchVision v2中:
ToImage()的作用是将输入转换为PIL Image或tensor格式,但不会处理像素值范围和数据类型;ToTensor()会将PIL Image/numpy数组转换为float32 tensor,并将像素值从[0,255]缩放到[0,1];v2.ToDtype(torch.float32, scale=True)和ToTensor()功能一致,且更符合v2的API风格,可替代ToTensor()。
推荐修改Transform为:
data_transform = v2.Compose([ v2.ToImage(), v2.RandomResizedCrop(size=(680, 680), antialias=True), v2.ToDtype(torch.float32, scale=True), # 替代ToTensor,更适配v2生态 ])
内容的提问来源于stack exchange,提问作者Conner Carriere
相关产品推荐
相关产品推荐

