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

如何在PyTorch自定义Dataset上应用自定义Transform数据变换

PyTorch自定义数据集Transform不生效问题解决方案

核心原因

  • 你打印的是数据集实例的samples属性,该属性是你初始化数据集时传入的原始numpy数组,Transform的变换逻辑仅在__getitem__方法中触发,只有通过索引访问数据集元素(比如transformed[0])、或者使用DataLoader加载数据时才会执行变换,直接访问原始存储的samples属性不会触发任何变换操作。
  • 你传入数据集的原始数据是numpy数组类型,但自定义的Transform都是基于PyTorch张量的操作,直接传入numpy数组会导致类型不匹配,需要先把numpy数组转成PyTorch张量再执行变换。
  • 测试代码中你定义的数据集类名为MyDataset,但实例化时用的是NumbersDataset,类名不一致会触发命名错误,属于笔误类问题。

修复后的可运行代码示例

import torch
import numpy as np
from torch.utils.data import Dataset
from torchvision import transforms

# 自定义高斯噪声变换
class AddGaussianNoise(object): 
    def __init__(self, mean, std):
        self.std = std
        self.mean = mean
        
    def __call__(self, tensor):
        return tensor + torch.randn(tensor.size()) * self.std + self.mean
    
    def __repr__(self):
        return self.__class__.__name__ + f'(mean={self.mean}, std={self.std})'   

# 自定义归一化变换
class Normalize(object):
    def __init__(self, mean, std):
        self.std = std
        self.mean = mean
        
    def __call__(self, tensor):
        return (tensor.sub_(self.mean)).div(self.std)
    
    def __repr__(self):
        return self.__class__.__name__ + f'(mean={self.mean}, std={self.std})'  

# 自定义数据集
class MyDataset(Dataset):
    def __init__(self, data, transforms = None):
        self.samples = data
        self.transforms= transforms
    def __len__(self):
        return len(self.samples)

    def __getitem__(self, idx):
        sample = self.samples[idx]
        # 先把numpy数组转换为float类型的张量,适配后续变换操作
        sample = torch.tensor(sample, dtype=torch.float32)
        if self.transforms is not None:
            sample = self.transforms(sample)
        return sample

# 测试逻辑
if __name__ == "__main__":
    data = np.array([[-1,-1,1,-1],[-1,1,-1,-1],[1,-1,-1,-1],[-1,-1,-1,1]])
    transformed_dataset = MyDataset(data, transforms.Compose([
        AddGaussianNoise(0.5, 0.5),
        Normalize(0.5, 0.5),
    ]))
    # 通过索引访问元素触发变换,查看效果
    for idx in range(len(transformed_dataset)):
        print(f"第{idx}个样本变换后结果:\n{transformed_dataset[idx]}\n")

补充说明

你单独调用Transform可以正常生效的原因是,你直接把完整数据传给了Transform实例主动执行了变换,没有经过数据集的懒加载逻辑,所以可以直接得到变换后的结果。如果你的需求是对整个二维数组做变换,而非按行拆分样本,修改__getitem__方法的索引逻辑即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 15:24:09