如何改写FiftyOne数据集类以适配PyTorch DataLoader?
解决FiftyOne与PyTorch DataLoader结合的索引访问问题
问题背景
想用公开数据集的边界框做初始训练,选择FiftyOne工具,但与PyTorch结合时遇到索引访问错误。官方示例已过时,FiftyOne的Dataset/DatasetView不支持数字索引(仅能通过dataset.first()、dataset.last()等方法访问样本),导致PyTorch DataLoader抛出以下异常:
KeyError: 'Accessing samples by numeric index is not supported. Use sample IDs, filepaths, slices, boolean arrays, or a boolean ViewExpression instead'
我修改了官方示例代码,仅保留汽车子集,但调用torch_dataset_test[0]时直接触发上述错误:
class FiftyOneDS(torch.utils.data.Dataset): def __init__( self , fiftyone_ds , transforms = None , gt_field = "ground_truth" , classes = None ): self.samples = fiftyone_ds self.transforms = transforms self.gt_field = gt_field self.classes = classes # don't care self.img_paths = self.samples.values("filepath") def __getitem__(self, idx): img_path = self.img_paths[idx] sample = self.samples[idx] # 此处触发错误:FiftyOne不支持数字索引 metadata = sample.metadata img = Image.open(img_path).convert("RGB") boxes = [] labels = [] detections = sample[self.gt_field].detections for det in detections: if det["label"] != "car": continue category_id = self.labels_map_rev[det.label] # 原代码未定义该变量 coco_obj = fouc.COCOObject.from_label( det, metadata, category_id=category_id, ) x, y, w, h = coco_obj.bbox boxes.append([x, y, x + w, y + h]) labels.append(coco_obj.category_id) target = {} target["boxes"] = torch.as_tensor(boxes, dtype=torch.float32) target["labels"] = torch.as_tensor(labels, dtype=torch.int64) target["image_id"] = torch.as_tensor([idx]) if self.transforms is not None: img, target = self.transforms(img, target) return img, target def __len__(self): return len(self.img_paths)
使用代码片段:
carset = FiftyOneDS(dataset) print("type:", type(carset)) # type: <class '__main__.FiftyOneDS'> print("first elem:", carset[0]) # 触发KeyError
解决方案
核心思路是预先提取FiftyOne样本的唯一标识(如样本ID)为Python原生列表,通过标识而非数字索引获取样本。以下是改写后的数据集类:
import torch from PIL import Image import fiftyone as fo import fiftyone.utils.coco as fouc class FiftyOneDS(torch.utils.data.Dataset): def __init__( self , fiftyone_ds , transforms = None , gt_field = "ground_truth" , classes = None ): self.samples = fiftyone_ds self.transforms = transforms self.gt_field = gt_field self.classes = classes # 提取样本ID列表(转为原生Python列表确保索引正常) self.sample_ids = list(self.samples.values("id")) # 预定义标签映射(根据你的数据集类别调整) self.labels_map_rev = {"car": 1} # 示例:将"car"映射为类别ID 1 def __getitem__(self, idx): # 通过样本ID获取FiftyOne样本 sample_id = self.sample_ids[idx] sample = self.samples[sample_id] img_path = sample.filepath metadata = sample.metadata img = Image.open(img_path).convert("RGB") boxes = [] labels = [] detections = sample[self.gt_field].detections for det in detections: if det.label != "car": continue category_id = self.labels_map_rev[det.label] coco_obj = fouc.COCOObject.from_label( det, metadata, category_id=category_id, ) x, y, w, h = coco_obj.bbox boxes.append([x, y, x + w, y + h]) labels.append(coco_obj.category_id) target = {} target["boxes"] = torch.as_tensor(boxes, dtype=torch.float32) target["labels"] = torch.as_tensor(labels, dtype=torch.int64) target["image_id"] = torch.as_tensor([idx]) if self.transforms is not None: img, target = self.transforms(img, target) return img, target def __len__(self): return len(self.sample_ids)
关键改动说明
- 预提取样本ID列表:将FiftyOne返回的
values("id")转为Python原生列表,确保可以通过数字索引获取样本ID - 通过样本ID访问样本:在
__getitem__中用self.samples[sample_id]替代原有的self.samples[idx],符合FiftyOne的访问规则 - 补充标签映射:修复原代码中未定义的
labels_map_rev变量,根据实际数据集类别调整映射关系
使用示例
# 假设dataset是你的FiftyOne数据集/视图 carset = FiftyOneDS(dataset) dataloader = torch.utils.data.DataLoader(carset, batch_size=4, shuffle=True) # 测试加载数据 for imgs, targets in dataloader: print(imgs.shape) print(targets["boxes"].shape) break
内容的提问来源于stack exchange,提问作者user3002166
相关产品推荐
相关产品推荐

