一、生成图像路径文档的函数 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. 类定义:MyDatasetDataset

继承 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)

三、定义预处理操作(通常放在数据集类之后,主逻辑之前)

  1. 统一数据格式:模型对输入图像的尺寸、通道数(如 RGB)、数据类型(如张量)有固定要求,预处理需将原始图像转换为符合要求的格式。
  2. 消除数据偏差:原始图像可能存在亮度、对比度差异,通过归一化等操作可让数据分布更稳定,帮助模型快速收敛。
  3. 增强数据多样性:通过随机裁剪、翻转等操作生成 “新样本”,减少模型对训练数据的过拟合,提升泛化能力。
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 批量加载

不适用的场景(需要修改框架)

  1. 非图像分类任务:如图像分割(需要掩码文件)、目标检测(需要标注框)等,需在MyDataset中添加对应标签的解析逻辑。
  2. 标签不基于文件夹名:若标签存在于单独的标注文件(如 CSV、XML)中,需修改classes_txtMyDataset的标签读取方式。
  3. 数据集结构不规则:如所有图像在同一文件夹,标签通过文件名后缀区分(如image_0_class1.jpg),需修改路径解析逻辑。
  4. 非 PyTorch 框架:若使用 TensorFlow,需将Dataset类改为tf.data.Dataset的实现方式。

总结

这个框架是 「按类别分文件夹的图像分类任务」在 PyTorch 中的标准实现模板 ,核心逻辑(路径索引生成 + 数据集加载)具有通用性,只需根据具体场景调整细节(如标签格式、预处理方式、文件路径分隔符等)即可复用。

更多推荐