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

PyTorch 如何为自定义数据集手动应用归一化变换供DataLoader使用

自定义PyTorch数据集应用transform预处理方案

要实现和内置数据集完全一致的预处理逻辑,你只需要继承PyTorch的Dataset抽象类实现自定义数据集,将transform逻辑嵌入到样本读取流程中即可,具体实现如下:

首先导入依赖:

import torch
import pandas as pd
from torchvision import transforms
from torch.utils.data import Dataset, DataLoader

然后实现自定义数据集类:

class SignMNISTDataset(Dataset):
    def __init__(self, csv_path, transform=None):
        self.data = pd.read_csv(csv_path)
        self.transform = transform

    def __len__(self):
        # 返回数据集总样本量
        return len(self.data)

    def __getitem__(self, idx):
        # 拆分单条样本的标签和像素数据
        label = self.data.iloc[idx, 0]
        # 784维像素值还原为28*28的单通道灰度图格式
        img = self.data.iloc[idx, 1:].values.reshape(28, 28).astype('uint8')
        
        # 应用传入的预处理规则
        if self.transform:
            img = self.transform(img)
        return img, label

之后就可以和教程中的内置数据集用法完全一致调用:

# 和教程保持完全相同的预处理规则
mean, std = (0.5,), (0.5,)
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize(mean, std)
])

# 实例化自定义训练集
trainset = SignMNISTDataset(
    csv_path='../input/sign-language-mnist/sign_mnist_train.csv',
    transform=transform
)
# 构造DataLoader
trainloader = DataLoader(trainset, batch_size=64, shuffle=True)

以上实现的预处理逻辑和教程完全等价:ToTensor()会把0-255范围的uint8像素值转换为0-1范围的浮点张量,再经过Normalize处理后缩放至-1到1范围,适配后续神经网络输入要求。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 17:36:03