如何修改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
相关产品推荐
相关产品推荐

