PyTorch中无需自定义Dataloader,如何用ImageFolder对接Albumentations做数据增强
错误原因分析
- 第一个报错:Albumentations默认输出HWC(高度、宽度、通道)格式的numpy数组,而PyTorch的卷积层要求输入为CHW(通道、高度、宽度)格式的张量,你第一次调用时没有加转张量的操作,模型将第二个维度32识别为输入通道数,因此报通道不匹配的错误。
- 第二个报错:
ToTensorV2仅会将numpy数组转为张量,不会自动将uint8类型的像素值(0255)归一化到01的float32类型,因此输入模型的张量为ByteTensor,和模型权重的FloatTensor类型不兼容。
完整适配实现
第一步:导入依赖
import numpy as np import albumentations as A from albumentations.pytorch import ToTensorV2 from torchvision import datasets
第二步:编写通用Transform包装类
这个类可以适配所有Albumentations的增强流水线,无需重复修改
class AlbumentationsTransform: def __init__(self, transform_pipeline): self.transform = transform_pipeline def __call__(self, img): # ImageFolder默认读取的是PIL Image,先转numpy数组 img_np = np.array(img) # 执行增强,返回处理后的张量 augmented = self.transform(image=img_np) return augmented['image']
第三步:定义增强流水线并调用ImageFolder
# 定义增强流水线,注意添加归一化操作 train_transform = A.Compose( [ A.Resize(height=224, width=224), # 可选添加其他增强,比如随机水平翻转、颜色抖动等 A.HorizontalFlip(p=0.5), A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), # 归一化到0~1并标准化,自动转float32 ToTensorV2(), # 转CHW格式的张量 ] ) # 直接传入ImageFolder即可 trainset = datasets.ImageFolder(root=traindir, transform=AlbumentationsTransform(train_transform))
补充说明
如果你不需要用ImageNet的标准化参数,也可以把A.Normalize替换为A.ToFloat(max_value=255.0),即可实现仅将像素值转为0~1的float32类型,满足类型匹配要求。
内容的提问来源于stack exchange,提问作者AJW
相关产品推荐
相关产品推荐

