深度学习图像分类实战指南之 数据预处理级
一、生成图像路径文档的函数 classes_txt(root,out_path,num_class=None)
1. 导入依赖库
os:用于文件 / 路径操作;torch/torch.nn等:PyTorch 核心库(后续模型会用到);Dataset:PyTorch 的数据集基类(自定义数据集需要继承它);PIL.Image:用于读取图像。
import os
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision.transforms as transforms
from torch.utils.data import DataLoader, Dataset
from PIL import Image
2. 函数定义:classes_txt(root,out_path,num_class=None)
root:数据集根目录(目录下是按类别分的子文件夹,每个子文件夹对应一个汉字类别);out_path:输出的 TXT 文件路径(用于保存所有图像的路径);num_class:要处理的类别数量(默认处理所有类别)。
3. 函数内部逻辑
函数会读取 root(数据集根目录)下的所有类别文件夹(如 0/、1/、2/ 等),然后遍历每个文件夹下的图像文件,将图像的完整路径(如 root/0/img1.jpg)收集起来。
(1).“扫描类别文件夹”:确定待处理的类别范围
def classes_txt(root, out_path, num_class=None):
dirs = os.listdir(root) # 获取所有类别文件夹名(如0、1、2...)
if not num_class:
num_class = len(dirs)
先读取数据集根目录下的所有子文件夹(每个文件夹对应一个类别,比如 “0”“1”“2” 分别代表不同汉字);
如果用户没指定num_class(要处理的类别数量),就默认处理所有类别。
(2).“初始化路径文档”:避免文件不存在的异常
# 若输出文件不存在,创建空文件
if not os.path.exists(out_path):
f = open(out_path, 'w')
f.close()
如果输出的 TXT 文件(用来存图像路径)还没创建,就先新建一个空文件,避免后续写入时抛 “文件不存在” 的错误。
(3). “续写逻辑”:只补充未处理的类别(核心亮点)
with open(out_path, 'r+') as f: #
# 尝试读取已有文件的最后一行,获取已处理的最后一个类别
end = int(f.readlines()[-1].split('/')[-2]) + 1
except:
end = 0 # 若文件为空或读取失败,从第0类开始
用try-except处理两种情况:
如果 TXT 文件已有内容:读取最后一行路径,从路径中提取 “上一次处理到的类别编号”(比如路径是root/5/img.jpg,则提取5),end设为5+1=6,表示下一次从第 6 类开始处理;
如果 TXT 文件为空:end设为0,表示从第 0 类开始处理。
这时候,可能有些同学会对with open ( )as f:还有 r+ 有些陌生,我在这里介绍一下
在 Python 中,with 是一种用于 “上下文管理” 的语法结构,主要作用是 “自动管理资源”(比如文件、网络连接等),确保资源在使用完毕后被正确释放,避免资源泄露。
在文件操作中,r+ 是 文件打开模式 的一种,表示 “读写模式”(read and write),具体含义和作用如下:
对 r+ 的拆解:
r:表示 只读模式(read),打开文件后可以读取内容,但默认不能写入(如果只写 r,写入会报错)。
+:表示 扩展模式,与 r 结合后(r+),在 “只读” 基础上增加了 “写入” 权限,即 既能读文件内容,又能向文件写入内容。
如果不用 with,实现同样的功能需要这样写:
#手动打开文件
f=open(out_path, 'r+')
try:
#读写操作
#... ...
finally:
#确保无论是否出错,文件都会被关闭
f.close()
显然,with 语法更简洁,且能避免人为遗漏 close() 的风险。
(4).“筛选未处理类别 + 写入路径”:完成路径生成
if end < num_class - 1:
dirs.sort() # 按类别文件夹名排序(确保0、1、2...顺序)
dirs = dirs[end:num_class] # 截取需要补充的类别
for dir in dirs: # 遍历每个类别文件夹
files = os.listdir(os.path.join(root, dir)) # 获取该类别下的所有图像
for file in files: # 遍历每张图像
# 拼接完整路径并写入 txt(如root/0/image1.jpg)
f.write(os.path.join(root, dir, file) + '\n')
先判断 “已处理的类别数” 是否小于 “需要处理的类别数 - 1”:如果是,说明还有类别没处理完;
对类别文件夹列表dirs排序(保证每次处理的类别顺序一致,避免混乱);
截取[end:num_class]的类别(只处理未生成过路径的类别);
遍历每个类别文件夹,再遍历文件夹内的所有图像,把完整图像路径写入 TXT 文件。
二、自定义数据集类(MyDataset)
作用是读取 TXT 中的图像路径,加载图像并转换成 PyTorch 张量,用于后续模型训练 / 测试。
1. 类定义:MyDataset(Dataset)
继承 PyTorch 的Dataset基类,必须实现__init__、__getitem__、__len__三个方法。
# 数据加载及预处理
class MyDataset(Dataset):
2. __init__方法(初始化)
def __init__(self, txt_path, num_class, transforms=None):
super().__init__() # 继承父类(Dataset)的初始化
# 存储图像路径
images = []
# 存储图像对应的类别标签(在本例中是汉字对应的数字ID)
labels = []
# 打开上一步生成的TXT文档
with open(txt_path, 'r') as f:
for line in f:
# 只读取前num_class个类别的图像(避免读取多余类别)
if int(line.split('\\')[-2]) >= num_class:
break
line = line.strip('\n') # 去掉换行符
images.append(line) # 保存图像路径
# 从路径中提取类别标签(文件夹名转成数字)
labels.append(int(line.split('\\')[-2]))
self.images = images # 类属性:所有图像路径
self.labels = labels # 类属性:所有图像的标签
self.transforms = transforms # 类属性:图像格式转换方法(如Resize、ToTensor等)
3. __getitem__方法(按索引取数据)
def __getitem__(self, index):
# 用PIL读取图像,并转成RGB格式
image = Image.open(self.images[index]).convert('RGB')
# 获取该图像对应的标签
label = self.labels[index]
# 如果指定了格式转换方法,就对图像进行转换(如转张量、归一化等)
if self.transforms is not None:
image = self.transforms(image)
# 返回图像张量和标签
return image, label
4. __len__方法(返回数据集长度)
def __len__(self):
return len(self.labels)
三、定义预处理操作(通常放在数据集类之后,主逻辑之前)
-
统一数据格式:模型对输入图像的尺寸、通道数(如 RGB)、数据类型(如张量)有固定要求,预处理需将原始图像转换为符合要求的格式。
-
消除数据偏差:原始图像可能存在亮度、对比度差异,通过归一化等操作可让数据分布更稳定,帮助模型快速收敛。
-
增强数据多样性:通过随机裁剪、翻转等操作生成 “新样本”,减少模型对训练数据的过拟合,提升泛化能力。
def get_transforms(is_train=True):
"""根据训练/测试阶段返回不同的预处理流程"""
if is_train:
# 训练集:添加数据增强
return transforms.Compose([
transforms.Resize((224, 224)), # 调整尺寸
transforms.RandomHorizontalFlip(p=0.5), # 随机水平翻转
transforms.RandomCrop(200), # 随机裁剪
transforms.ToTensor(), # 转为张量
transforms.Normalize(mean=[0.485, 0.456, 0.406], # ImageNet均值
std=[0.229, 0.224, 0.225]) # ImageNet标准差
])
else:
# 测试集:只做必要的预处理(无数据增强)
return transforms.Compose([
transforms.Resize((200, 200)), # 直接调整到目标尺寸(与训练集裁剪后一致)
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225])
])
四、运用场景
这段代码的框架适用于有明确类别划分的图像分类任务,尤其是当数据集满足以下特征时,基本可以直接复用或稍作修改:
1. 数据集目录结构符合「按类别分文件夹」的规范
例如:
- 手写数字识别(MNIST 的自定义版本,按 0-9 分文件夹)
- 花卉分类(玫瑰、百合等按类别分文件夹)
- 目标分类(猫、狗、汽车等按类别分文件夹)
2. 需要生成「图像路径 + 标签」的索引文件(txt)
当你需要以下功能时,classes_txt函数的逻辑可以直接复用:
- 避免每次加载数据时重复遍历文件夹(通过 txt 文件快速索引)
- 灵活控制使用的类别数量(通过
num_class参数选择前 N 类) - 支持增量补充数据(原逻辑意图:若 txt 文件已有部分类别,可补充剩余类别,修复写入模式后即可实现)
3. 基于 PyTorch 的自定义数据集加载
当你使用 PyTorch 框架,且需要对图像进行预处理(如 resize、归一化、数据增强等)时,MyDataset类的框架适用:
- 需要将图像路径和标签对应起来
- 需要在加载时动态应用预处理(通过 transforms 参数)
- 需要符合 PyTorch 的 Dataset 规范,方便后续用 DataLoader 批量加载
不适用的场景(需要修改框架)
- 非图像分类任务:如图像分割(需要掩码文件)、目标检测(需要标注框)等,需在
MyDataset中添加对应标签的解析逻辑。 - 标签不基于文件夹名:若标签存在于单独的标注文件(如 CSV、XML)中,需修改
classes_txt和MyDataset的标签读取方式。 - 数据集结构不规则:如所有图像在同一文件夹,标签通过文件名后缀区分(如
image_0_class1.jpg),需修改路径解析逻辑。 - 非 PyTorch 框架:若使用 TensorFlow,需将
Dataset类改为tf.data.Dataset的实现方式。
总结
这个框架是 「按类别分文件夹的图像分类任务」在 PyTorch 中的标准实现模板 ,核心逻辑(路径索引生成 + 数据集加载)具有通用性,只需根据具体场景调整细节(如标签格式、预处理方式、文件路径分隔符等)即可复用。
更多推荐
所有评论(0)