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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 09:45:01