如何基于带标签自定义数据集构建PyTorch数据类训练分类器
解决方案
1. 创建分类文件夹结构
先确保目标目录和分类子文件夹存在,用os模块实现:
import os base_dir = "/content/sub_f" # 创建基础目录,已存在则不报错 os.makedirs(base_dir, exist_ok=True) # 循环创建4个分类子文件夹 for class_idx in range(4): class_dir = os.path.join(base_dir, str(class_idx)) os.makedirs(class_dir, exist_ok=True)
2. 遍历Dataloader并将图片存入对应分类文件夹
需要把Tensor格式的图片转换成PIL图像后保存,同时处理Dataloader返回的批量数据:
from PIL import Image import torchvision.transforms as transforms # 定义Tensor转PIL图像的工具 to_pil = transforms.ToPILImage() # 遍历整个Dataloader的数据 for batch_idx, (imgs, labels) in enumerate(dataloader): # 处理当前批次里的每张图片和对应标签 for img_idx, (img, label) in enumerate(zip(imgs, labels)): class_idx = label.item() # 生成唯一的保存路径,避免重名 save_path = os.path.join(base_dir, str(class_idx), f"img_{batch_idx}_{img_idx}.png") # 转换格式并保存 pil_img = to_pil(img) pil_img.save(save_path)
提示:如果Dataloader返回的标签张量维度特殊,可根据实际情况调整取值方式,比如
label.squeeze().item()。
3. 使用ImageFolder加载数据集
定义预处理转换后,直接加载数据集并验证:
from torchvision.datasets import ImageFolder import torchvision.transforms as transforms # 定义数据预处理组合(可根据训练需求调整) trans_comp = transforms.Compose([ transforms.Resize((256, 256)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 加载数据集 dataset = ImageFolder(root=base_dir, transform=trans_comp) # 验证分类是否正确 print(dataset.classes) # 预期输出 ['0', '1', '2', '3'] print(dataset.class_to_idx) # 查看分类与索引的映射关系 # 生成训练用的Dataloader train_loader = torch.utils.data.DataLoader(dataset, batch_size=32, shuffle=True)
额外提示
- 10000张图片的保存操作耗时较长,建议在GPU环境下执行,或者分批次处理。
- 如果原始图片本身已存储在磁盘中,直接复制文件到对应分类文件夹会比Tensor转存效率更高。
内容的提问来源于stack exchange,提问作者Formal_this
相关产品推荐
相关产品推荐

