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

如何改写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)

关键改动说明

  1. 预提取样本ID列表:将FiftyOne返回的values("id")转为Python原生列表,确保可以通过数字索引获取样本ID
  2. 通过样本ID访问样本:在__getitem__中用self.samples[sample_id]替代原有的self.samples[idx],符合FiftyOne的访问规则
  3. 补充标签映射:修复原代码中未定义的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 12:06:08