如何在DataLoader外使用transforms.FiveCrop()并保留批量维度
解决批量图像变换后维度缺失的问题
你的核心问题是没有将每个样本的crop结果合并为一个包含batch维度的张量。只需在函数最后对结果列表执行torch.stack()操作,就能得到期望的[bs, ncrops, c, h, w]形状。同时可以优化变换逻辑,让代码更简洁:
import torch from torchvision import transforms as TF def get_patches(orig_img): # 保存原数据所在设备,确保结果返回对应设备 device = orig_img.device # 将批量张量拆分为单张PIL图像 images = [TF.to_pil_image(x) for x in orig_img.cpu()] # 定义单张crop的后续变换组合 post_crop_transform = TF.Compose([ TF.Resize(128), TF.ToTensor(), TF.Normalize([0.5], [0.5]) ]) resized_imgs = [] for img in images: # 执行CenterCrop img_cropped = TF.CenterCrop(100)(img) # 执行FiveCrop,得到5张裁剪图 five_crops = TF.FiveCrop(64)(img_cropped) # 注意:你问题描述中是16,代码里写的是64,请按需调整 # 对每个crop应用变换并堆叠为[ncrops, c, h, w] crop_tensor = torch.stack([post_crop_transform(crop) for crop in five_crops]) resized_imgs.append(crop_tensor) # 将所有样本的crop结果堆叠,得到[bs, ncrops, c, h, w],并转回原设备 return torch.stack(resized_imgs).to(device) # 使用示例 orig_img = next(iter(DataLoader)) patches = get_patches(orig_img) print(patches.shape) # 输出: torch.Size([4, 5, 1, 128, 128])
关键说明
torch.stack(resized_imgs):将列表中每个形状为[5, 1, 128, 128]的张量在第0维堆叠,直接生成包含batch维度的5维张量。TF.Compose:把重复的变换逻辑封装成组合操作,提升代码可读性和可维护性。- 设备同步:避免处理后的数据留在CPU,确保和原数据设备一致(GPU/CPU)。
- 参数核对:注意
FiveCrop的参数,你问题描述中是16,但代码里写的是64,根据实际需求调整即可。
内容的提问来源于stack exchange,提问作者Ammar Ul Hassan
相关产品推荐
相关产品推荐

