torchvision.datasets.ImageFolder 专门用来读取按文件夹分类的图像数据集,是图像分类任务最常用的自定义数据集类。
一、数据集目录强制格式
必须遵循下面的层级:
root/
类别A/
图片1.jpg
图片2.png
类别B/
图片3.jpg
...
root:数据集根路径(传给第一个参数)- 一级子文件夹名称 = 类别名
- 子文件夹里存放该类别的所有图片
- ImageFolder 会自动给每个类别分配数字标签(0,1,2...)
例如:猫狗分类文件结构

规则:只能一层类别文件夹,不能多层嵌套。
二、函数原型
python
ImageFolder(
root,
transform=None,
target_transform=None,
loader=default_loader,
is_valid_file=None
)
- root(str) 数据集根目录路径。
- transform (callable, 可选) 对图像本身做预处理、数据增强。 接收 PIL 图片,返回处理后的图片 / Tensor。 示例:Resize、随机翻转、转 Tensor、归一化。
- target_transform (callable, 可选) 对标签做变换。 比如把数字标签转成 one‑hot 编码。
- loader 图片读取函数,默认
default_loader,用 PIL 读取图片。一般不用修改。 - is_valid_file 过滤文件,自定义哪些文件才视为有效图片。
三、返回对象的结构
实例化之后得到一个 Dataset 对象:
- 遍历单条样本:
(image, label)- image:经过 transform 后的图像张量 / PIL 图
- label:该类别的数字索引(int)
额外自带两个重要属性:
dataset.classes→ 类别名称列表['cat','dog']dataset.class_to_idx→ 类别→数字映射字典{'cat':0, 'dog':1}
四、运行示例
python
from torchvision import datasets, transforms
trans = transforms.Compose([
transforms.Resize((224,224)),
transforms.ToTensor()
])
dataset = datasets.ImageFolder(root="./train", transform=trans)
print(dataset.classes)
print(dataset.class_to_idx)
img, label = dataset[0]
print(img.shape, label)
五、搭配 DataLoader 使用(训练标准写法)
python
import torch
dataloader = torch.utils.data.DataLoader(
dataset,
batch_size=batch_size,
shuffle=true,
num_workers=num_workers,
pin_memory=True,
drop_last=False,
)
六、调用链路示意图

流程:
ImageFolder扫描文件夹,自动生成图片路径 + 标签transform流水线对每张图片做缩放、翻转、转张量、归一化- 交给
DataLoader打包成 batch,送入网络训练
七、总结
ImageFolder 自动扫描指定目录下的子文件夹,读取图片并生成分类数据集,配合 transform 做预处理,用于 PyTorch 图像分类训练。