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
相关产品推荐
相关产品推荐

