PyTorch中如何加载.npy格式的输入输出文件夹配对数据
解决方案
原生DatasetFolder的设计逻辑是针对「单根目录下不同子文件夹对应不同类别标签」的单输入分类场景,不直接支持跨两个独立文件夹找(X,Y)配对文件的需求,硬改它的内部方法反而冗余,直接继承Dataset基类写自定义数据集是最简单、可控的方案,十几行代码就能搭完可用的版本。
第一步:确认目录结构
确保你的文件结构符合如下形式,配对文件文件名完全一致:
your_project/ └── dataset/ ├── input/ │ ├── 001.npy │ ├── 002.npy │ └── ...(所有输入npy) └── output/ ├── 001.npy ├── 002.npy └── ...(所有对应输出npy,和input下同名文件一一配对)
第二步:编写自定义配对数据集类
代码内置了配对校验逻辑,会自动过滤两边文件夹里不匹配的孤立文件:
import os import numpy as np import torch from torch.utils.data import Dataset, random_split, DataLoader class PairedNpyDataset(Dataset): def __init__(self, input_dir, output_dir, transform=None): self.input_dir = input_dir self.output_dir = output_dir self.transform = transform # 提取两个文件夹下所有npy文件名,取交集得到有效配对 input_file_set = {f for f in os.listdir(input_dir) if f.endswith(".npy")} output_file_set = {f for f in os.listdir(output_dir) if f.endswith(".npy")} self.paired_files = sorted(list(input_file_set & output_file_set)) # 基础校验,避免路径写错或者文件缺失 if len(self.paired_files) == 0: raise RuntimeError("未找到任何匹配的输入输出npy文件对,请检查文件夹路径是否正确") if len(input_file_set) != len(self.paired_files) or len(output_file_set) != len(self.paired_files): print(f"提示:存在未配对的孤立文件,最终加载有效样本对共{len(self.paired_files)}组") def __len__(self): return len(self.paired_files) def __getitem__(self, index): file_name = self.paired_files[index] # 加载对应配对的输入、输出数组 X = np.load(os.path.join(self.input_dir, file_name)) Y = np.load(os.path.join(self.output_dir, file_name)) # 默认把numpy数组转成PyTorch浮点张量,可自行添加归一化、维度调整逻辑 X = torch.from_numpy(X).float() Y = torch.from_numpy(Y).float() # 应用自定义预处理/数据增强 if self.transform is not None: X = self.transform(X) Y = self.transform(Y) return X, Y
第三步:拆分数据集、构建训练用DataLoader
初始化数据集后直接用PyTorch内置的random_split拆分训练/测试集即可,不需要额外写拆分逻辑:
# 初始化全量数据集 full_dataset = PairedNpyDataset( input_dir="./dataset/input", output_dir="./dataset/output" ) # 按8:2比例拆分训练集、测试集 train_sample_count = int(0.8 * len(full_dataset)) test_sample_count = len(full_dataset) - train_sample_count train_dataset, test_dataset = random_split( full_dataset, [train_sample_count, test_sample_count], generator=torch.Generator().manual_seed(42) # 固定随机种子,保证拆分结果可复现 ) # 构建可直接送入训练循环的DataLoader train_loader = DataLoader(train_dataset, batch_size=16, shuffle=True, num_workers=2) test_loader = DataLoader(test_dataset, batch_size=16, shuffle=False, num_workers=2)
实用注意事项
- 如果你的npy文件单个体积很大,可以给
np.load加上mmap_mode="r"参数,用内存映射方式加载,避免一次性把全量数据读入内存占满显存/内存。 - 如果输入输出数组维度不符合PyTorch要求(比如图像类数据是
(H,W,C)的numpy格式,需要转成(C,H,W)的张量格式),可以直接在__getitem__方法里加维度转换逻辑,比如X = X.permute(2, 0, 1)。 - 不建议硬改
DatasetFolder适配这个场景:DatasetFolder内部默认是遍历单个根目录下的子文件夹打类别标签,适配跨文件夹配对需要重写它的文件扫描逻辑,工作量比直接写自定义Dataset更大,后续维护也麻烦。
内容的提问来源于stack exchange,提问作者Dulat Yussupaliyev
相关产品推荐
相关产品推荐

