如何在PyTorch中不使用ImageFolder构建PNG、TIF格式图像自定义数据集
PyTorch自定义数据集实现(无需ImageFolder,支持tif/png格式)
完全可以实现,PyTorch提供的Dataset基类就是用于自定义数据集的,完全不需要依赖ImageFolder的固定目录结构要求,对png、tif等任意图像格式都可以适配。
前置依赖
- 图像读取:普通tif/png直接用
Pillow即可,若为多通道、高bit深度的专业tif文件,推荐额外安装tifffile库 - 核心依赖:
torch、torchvision
自定义数据集代码示例
以下代码适配tif/png格式图像,支持自定义标签匹配规则,可根据自身需求调整逻辑:
import os import torch from torch.utils.data import Dataset from PIL import Image # 处理特殊tif时替换为 import tifffile as tiff class CustomImageDataset(Dataset): def __init__(self, img_dir, label_csv_path=None, transform=None, target_transform=None): """ 参数说明: img_dir: 所有图像存放的文件夹路径 label_csv_path: 标签csv文件路径,无标签场景(如自监督训练)可设为None transform: 图像预处理操作 target_transform: 标签预处理操作 """ self.img_dir = img_dir # 自动过滤文件夹内的png、tif、tiff格式文件,排除无关文件干扰 self.img_paths = [ os.path.join(img_dir, f) for f in os.listdir(img_dir) if f.lower().endswith(('.png', '.tif', '.tiff')) ] self.transform = transform self.target_transform = target_transform # 标签读取逻辑,示例适配csv格式:第一列为图像文件名,第二列为整数标签,可自行修改规则 self.label_map = {} if label_csv_path is not None: with open(label_csv_path, 'r', encoding='utf-8') as f: lines = f.readlines()[1:] # 跳过csv表头行 for line in lines: fname, label = line.strip().split(',') self.label_map[fname] = int(label) def __len__(self): # 返回数据集总样本量 return len(self.img_paths) def __getitem__(self, idx): # 读取单条样本 img_path = self.img_paths[idx] fname = os.path.basename(img_path) # 读取图像,特殊tif替换为 image = tiff.imread(img_path) 即可 image = Image.open(img_path).convert('RGB') # 单通道灰度图改为 'L' # 读取对应标签,无标签场景可删除这段逻辑 label = self.label_map.get(fname, 0) # 应用预处理规则 if self.transform: image = self.transform(image) if self.target_transform: label = self.target_transform(label) return image, label
调用示例
from torchvision import transforms from torch.utils.data import DataLoader # 定义图像预处理逻辑 transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 实例化自定义数据集 dataset = CustomImageDataset( img_dir="./your_tif_images_folder", # 替换为你的tif图像文件夹路径 label_csv_path="./your_labels.csv", # 无标签时删除该参数 transform=transform ) # 封装为DataLoader即可直接用于训练 dataloader = DataLoader(dataset, batch_size=32, shuffle=True)
内容的提问来源于stack exchange,提问作者beginner
相关产品推荐
相关产品推荐

