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

如何修改PyTorch Dataset的__getitem__方法返回含10张图像的图像包

实现方案

核心思路是自定义继承ImageFolder的数据集类,将单张样本的索引逻辑重构为固定10张图为一个样本包,不需要改动原有数据存储结构。

  • 按类别(子文件夹)对图像索引分组,避免不同类别的图像被分到同一个包中
  • 每个类别下按文件读取顺序切分固定大小的图像包,末尾不足10张的残段默认丢弃,保证所有输出包大小一致
  • DataLoaderbatch_size设为1时,输出张量形状直接匹配需求

可直接运行的代码
import torch
from torchvision import datasets, transforms
from torch.utils.data import DataLoader

class BagImageFolder(datasets.ImageFolder):
    def __init__(self, root, transform=None, bag_size=10):
        super().__init__(root, transform=transform)
        self.bag_size = bag_size
        self.bag_list = []

        # 按类别归集样本索引
        class_idx_map = {}
        for sample_idx, (_, label) in enumerate(self.samples):
            if label not in class_idx_map:
                class_idx_map[label] = []
            class_idx_map[label].append(sample_idx)
        
        # 逐类别切分图像包
        for idx_list in class_idx_map.values():
            # 步长为包大小,丢弃末尾不足10张的部分
            for start in range(0, len(idx_list) - self.bag_size + 1, self.bag_size):
                self.bag_list.append(idx_list[start:start+self.bag_size])

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

    def __getitem__(self, item):
        sample_ids = self.bag_list[item]
        img_tensor_list = []
        label_list = []
        for sid in sample_ids:
            img_path, label = self.samples[sid]
            img = self.loader(img_path)
            if self.transform:
                img = self.transform(img)
            img_tensor_list.append(img)
            label_list.append(label)
        # 单包输出形状为[10, 3, 256, 256]
        return torch.stack(img_tensor_list, dim=0), torch.tensor(label_list)


# 调整transform参数匹配目标尺寸
transform = transforms.Compose([
    transforms.Resize(255),
    transforms.CenterCrop(256), # 原代码是224,改为256匹配目标输出尺寸
    transforms.ToTensor()
])

dataset = BagImageFolder('./../BCNB/patches/WSI_1', transform=transform, bag_size=10)
# batch_size=1时最终输出形状为[1, 10, 3, 256, 256]
data_loader = DataLoader(dataset, batch_size=1, shuffle=False)

# 形状验证
for batch_imgs, batch_labels in data_loader:
    print(batch_imgs.shape) # 输出 torch.Size([1, 10, 3, 256, 256])
    break

补充说明
  • 需要打乱顺序时直接将DataLoader的shuffle参数设为True即可,打乱单位是图像包,不会破坏单个包内的图像组成
  • 如果需要保留末尾不足10张的残包,可以修改切分bag_list的逻辑,对残包采用补零、重复采样等方式填充到10张
  • 包内图像顺序和原生ImageFolder读取顺序一致,默认按子文件夹内文件名排序读取

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 04:36:21