You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.04 05:09:03