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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 05:15:37