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

如何将Numpy ndarray图像数组转换为PyTorch数据集并使用DataLoader?

使用PyTorch Dataset和DataLoader实现批量图像处理

1. 自定义Dataset类

首先需要实现一个继承自torch.utils.data.Dataset的自定义类,负责单样本的格式转换逻辑:

import torch
from torch.utils.data import Dataset, DataLoader

class CustomImageDataset(Dataset):
    def __init__(self, images_np, labels_np):
        self.images = images_np
        self.labels = labels_np

    def __len__(self):
        return len(self.images)

    def __getitem__(self, idx):
        # 将单张numpy图像转为Tensor,并添加通道维度(变为1,128,128)
        img_tensor = torch.from_numpy(self.images[idx]).unsqueeze(0)
        # 转换标签为Tensor
        label_tensor = torch.from_numpy(self.labels[idx])
        return img_tensor, label_tensor

2. 实例化Dataset与DataLoader

用你的numpy数组初始化Dataset,再通过DataLoader实现批量处理、数据打乱、多进程加载等功能:

# 假设X是形状为(16699,128,128)的图像数组,y是对应标签数组
dataset = CustomImageDataset(X, y)
# 配置参数:batch_size根据显存调整,shuffle训练时建议开启,num_workers加速数据加载
dataloader = DataLoader(dataset, batch_size=32, shuffle=True, num_workers=4)

3. 简化训练循环

替换原来的逐样本循环,直接遍历DataLoader即可获取批量数据:

epochs = 3

for epoch in range(epochs):
    # 自动获取批量图像和标签,无需手动索引
    for batch_imgs, batch_labels in dataloader:
        # batch_imgs形状为(batch_size, 1, 128, 128),符合神经网络输入格式
        # 后续执行你的训练逻辑(前向传播、损失计算、反向传播等)
        print(f"Epoch {epoch}, Batch shape: {batch_imgs.shape}")

关键说明

  • Dataset类专注于单样本的格式转换,把numpy数组转为符合要求的Tensor格式
  • DataLoader自动完成批量拼接、数据打乱、多进程预加载,彻底替代手动循环索引的低效方式
  • 可根据硬件情况调整batch_size和num_workers参数,平衡显存占用和加载速度

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 09:07:11