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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 17:27:28