PyTorch DataLoader遍历Dataset时多目标加载格式异常问题
问题:DataLoader返回的targets格式不符合预期
遍历DataLoader时(for images, targets in dataloader),实际得到的targets格式为字典套列表:
{ 'image_id':[all image ids], 'keypoints':[all keypoint lists], 'labels':[all label lists], 'boxes':[all bboxes] }
但预期格式应为每个样本对应一个字典的列表:
[ { 'image_id':image_id of first sample, 'keypoints':[list of keypoints of first sample], 'labels':[list of labels of first sample], 'boxes':bbox of first sample }, ... { 'image_id':image_id of fourth sample, 'keypoints':[list of keypoints of fourth sample], 'labels':[list of labels of fourth sample], 'boxes':bbox of fourth sample } ]
当前数据集的__getitem__实现如下:
def __getitem__(self, idx): annotation = self.annotations[idx] image_id = annotation['image_id'] file_name = annotation['file_name'] image_path = f"{self.images_dir}/{file_name}" image = Image.open(image_path).convert("RGB") bbox = np.array(annotation['bbox']) keypoints = np.array([[ann["x"],ann["y"]] for ann in annotation["keypoints"]]) labels = np.array([kp_num[ann["name"]] for ann in annotation["keypoints"]]) target = { "image_id":image_id, "keypoints": torch.tensor(keypoints, dtype=torch.float32), "labels": torch.tensor(labels, dtype=torch.int64), "boxes":torch.tensor(bbox, dtype=torch.int) } if self.transform: image = self.transform(image) return image, target
曾尝试将target放入单元素列表返回,但结果仍是包含所有信息的单个条目,而非batch_size个目标的列表。使用的是torch.utils.data.DataLoader类。
解决方法
实现自定义collate_fn,手动整理batch中的图像和目标:
def custom_collate_fn(batch): images = [item[0] for item in batch] targets = [item[1] for item in batch] return images, targets
创建DataLoader时指定该函数即可:
dataloader = torch.utils.data.DataLoader( your_dataset, batch_size=4, collate_fn=custom_collate_fn )
内容的提问来源于stack exchange,提问作者fzLuHaef
相关产品推荐
相关产品推荐

