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

如何修改PyTorch的ImageFolder返回指定形状的图像张量

扩展ImageFolder实现10图组成的Bag张量返回

实现逻辑

原生ImageFolder已经做好了按类别文件夹解析样本、单图加载、transform流水线的能力,不需要重写底层逻辑,只需要在其基础上扩展同类别采样和张量堆叠的逻辑即可:

  • 初始化阶段额外维护一个按类别分组的索引字典,把每个类别对应的所有样本索引提前存好,避免每次采样遍历全量数据
  • 重写__getitem__方法,拿到当前样本的标签后,从对应类别的索引池里采样凑够10张图,逐张走transform后堆叠为4维张量
  • 自动适配类别样本数不足10的边界场景,采用有放回采样保证不会抛出索引错误
  • 最终输出的单样本形状为[10, 3, 256, 256],搭配batch_size=1的DataLoader即可得到需要的[1, 10, 3, 256, 256]形状,不需要额外维度调整

完整实现代码

import random
import torch
from torchvision.datasets import ImageFolder


class BagImageFolder(ImageFolder):
    def __init__(self, root, bag_size=10, transform=None, target_transform=None):
        # 复用原生ImageFolder的所有初始化逻辑
        super().__init__(
            root=root,
            transform=transform,
            target_transform=target_transform
        )
        self.bag_size = bag_size
        # 构建类别到样本索引的映射表
        self.class_to_indices = {}
        for idx, (_, label) in enumerate(self.samples):
            if label not in self.class_to_indices:
                self.class_to_indices[label] = []
            self.class_to_indices[label].append(idx)

    def __getitem__(self, index):
        # 获取当前样本的路径和标签
        path, target = self.samples[index]
        candidate_indices = self.class_to_indices[target]

        # 采样同类别下剩余的bag_size-1张图像
        if len(candidate_indices) >= self.bag_size:
            # 样本量充足时无放回采样,排除当前索引避免重复取图
            other_indices = random.sample(
                [i for i in candidate_indices if i != index],
                self.bag_size - 1
            )
        else:
            # 类别样本数不足bag_size时,有放回采样凑数
            other_indices = random.choices(candidate_indices, k=self.bag_size - 1)
        
        # 合并索引得到整个bag的样本列表
        bag_indices = [index] + other_indices
        bag_imgs = []
        for i in bag_indices:
            img_path, _ = self.samples[i]
            img = self.loader(img_path)
            if self.transform is not None:
                img = self.transform(img)
            bag_imgs.append(img)
        
        # 沿第0维堆叠,得到[bag_size, C, H, W]形状的张量
        bag_tensor = torch.stack(bag_imgs, dim=0)

        if self.target_transform is not None:
            target = self.target_transform(target)
        
        return bag_tensor, target

调用示例

from torchvision import transforms
from torch.utils.data import DataLoader

# 定义图像预处理流水线,和原生ImageFolder用法完全一致
transform = transforms.Compose([
    transforms.Resize((256, 256)),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

# 初始化数据集,指定数据集根目录、bag大小和预处理
dataset = BagImageFolder(
    root="./your_classification_dataset_root",
    bag_size=10,
    transform=transform
)
# DataLoader设置batch_size=1,即可得到目标形状的输出
dataloader = DataLoader(
    dataset,
    batch_size=1,
    shuffle=True,
    num_workers=4,
    pin_memory=True
)

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

可调说明

  • 如果不需要同类别采样(允许bag内出现不同类别的图像),可以删掉class_to_indices的构建逻辑,采样时直接从全量索引range(len(self.samples))中抽取即可
  • 如果需要固定bag的组成(每次运行同一个索引返回的10张图完全一致),把随机采样逻辑替换为固定偏移取图即可,比如按索引顺序取连续10张同类别图,边界处做取模处理
  • 如果需要复现实验结果,在训练脚本开头固定随机种子即可:random.seed(42)、torch.manual_seed(42)
  • 如果需要返回bag内每张图的单独标签,可以在堆叠图像的时候同步把每个索引对应的标签存成列表,和bag_tensor一起返回即可

内容的提问来源于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.29 09:24:13