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

PyTorch自定义DataLoader求助:特定目录结构数据集处理

自定义PyTorch DataLoader处理嵌套目录的Numpy数据集

针对你这种嵌套目录结构的数据集,核心思路是先提前收集所有数据与掩码的对应路径,再通过__getitem__方法按需加载单个样本。下面直接上可运行的代码和关键步骤解释:

步骤1:编写自定义Dataset类

PyTorch的DataLoader依赖Dataset类提供数据,我们需要继承torch.utils.data.Dataset并实现两个核心方法:__len__(返回总样本数)和__getitem__(返回指定索引的样本)。

import os
import numpy as np
import torch
from torch.utils.data import Dataset, DataLoader

class CustomNumpyDataset(Dataset):
    def __init__(self, root_dir):
        self.root_dir = root_dir
        self.data_pairs = []  # 存储(dat_path, mask_path)的对应关系
        
        # 遍历Numpy_dat下的所有子目录
        dat_root = os.path.join(root_dir, "Numpy_dat")
        mask_root = os.path.join(root_dir, "mask_numpy")
        
        for dat_subdir in os.listdir(dat_root):
            # 匹配对应的mask子目录(比如dat_0对应mask_0)
            mask_subdir = dat_subdir.replace("dat_", "mask_")
            dat_subdir_path = os.path.join(dat_root, dat_subdir)
            mask_subdir_path = os.path.join(mask_root, mask_subdir)
            
            if not os.path.exists(mask_subdir_path):
                continue  # 如果对应mask目录不存在,跳过
            
            # 遍历当前dat子目录下的所有npy文件
            for dat_file in os.listdir(dat_subdir_path):
                if not dat_file.endswith(".npy"):
                    continue
                # 匹配对应的mask文件名(比如dat_{0}_{0}.npy对应mask_{0}_{0}.npy)
                mask_file = dat_file.replace("dat_", "mask_")
                dat_path = os.path.join(dat_subdir_path, dat_file)
                mask_path = os.path.join(mask_subdir_path, mask_file)
                
                if os.path.exists(mask_path):
                    self.data_pairs.append((dat_path, mask_path))

    def __len__(self):
        # 返回总样本数
        return len(self.data_pairs)

    def __getitem__(self, idx):
        # 核心:根据索引加载单个样本
        dat_path, mask_path = self.data_pairs[idx]
        
        # 加载numpy数组
        dat_np = np.load(dat_path)
        mask_np = np.load(mask_path)
        
        # 转换为PyTorch张量(根据你的需求调整 dtype,比如float32)
        dat_tensor = torch.from_numpy(dat_np).float()
        mask_tensor = torch.from_numpy(mask_np).long()  # 如果是分类掩码,用long类型
        
        return dat_tensor, mask_tensor

步骤2:实例化Dataset和DataLoader

# 替换成你的All_data目录路径
dataset = CustomNumpyDataset(root_dir="./All_data")
dataloader = DataLoader(
    dataset,
    batch_size=4,  # 每次加载的样本数
    shuffle=True,  # 是否打乱数据
    num_workers=2  # 多进程加载(根据CPU核心数调整)
)

# 测试加载
for batch_dat, batch_mask in dataloader:
    print(f"数据批次形状: {batch_dat.shape}")
    print(f"掩码批次形状: {batch_mask.shape}")
    break

关键部分解释

  • __init__方法:提前遍历所有目录,把每个数据文件和对应的掩码文件路径配对存起来,避免每次加载数据时重复遍历目录,提升效率。
  • __getitem__方法:DataLoader会在需要样本时调用这个方法,传入索引idx,我们只需要根据索引取出对应的路径,加载numpy数组并转成张量即可。这一步是按需加载,不会一次性把所有数据塞进内存,适合大数据集。
  • 数据类型转换:根据你的任务调整,比如回归任务用float32,分类掩码用long(对应CrossEntropyLoss的要求)。

内容的提问来源于stack exchange,提问作者Uqhah

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 19:42:23