PyTorch修改图像像素无效问题求助:张量元素未改变
问题分析与解决办法
问题根源
每次访问trainset[0][0]时,都会重新从原CIFAR10数据集加载图片并执行完整的transform流程,返回的是新的临时张量。你修改的只是这个临时张量,并没有改变磁盘上的原始图片或数据集的底层缓存。再加上torch.utils.data.Subset本身只是原数据集的索引映射,不存储实际数据,每次取数都会调用原数据集的__getitem__方法重新生成张量,所以看起来修改完全没生效。
解决方法
方法1:将数据预加载到内存(适合CIFAR10这类小数据集)
把筛选后的所有数据一次性加载到内存中,这样修改后能永久保存:
transform = transforms.Compose( [transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))]) # 加载完整数据集 full_trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform) # 筛选目标类别并预加载数据到内存 indices_train = [i for i in range(len(full_trainset)) if full_trainset.targets[i] in [1,2]] data = [] targets = [] for idx in indices_train: img, label = full_trainset[idx] data.append(img) targets.append(label) # 自定义内存数据集类 class InMemoryDataset(torch.utils.data.Dataset): def __init__(self, data, targets): self.data = data self.targets = targets def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx], self.targets[idx] # 构建可修改的数据集 trainset = InMemoryDataset(data, targets) trainloader = torch.utils.data.DataLoader(trainset, batch_size=4, shuffle=True, num_workers=2) # 现在修改会永久生效 print(trainset[0][0][:,1,1]) trainset[0][0][:,1,1] = 0 print(trainset[0][0][:,1,1])
方法2:自定义Transform实现数据增强(对比学习常规做法)
对比学习的数据增强通常通过自定义Transform类实现,而非手动修改单个样本。比如要随机将某个像素置0,可以这样写:
import random class SetRandomPixelToZero: def __call__(self, img): # img是ToTensor处理后的张量 h, w = img.shape[1], img.shape[2] x = random.randint(0, h-1) y = random.randint(0, w-1) img[:, x, y] = 0.0 return img # 更新transform链 transform = transforms.Compose( [transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), SetRandomPixelToZero()]) # 后续数据集加载流程不变,每次取数自动应用增强 trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform) indices_train = [i for i in range(len(trainset)) if trainset.targets[i] in [1,2]] trainset = torch.utils.data.Subset(trainset, indices_train) trainloader = torch.utils.data.DataLoader(trainset, batch_size=4, shuffle=True, num_workers=2)
方法3:临时修改单一样本(仅用于测试)
如果只是临时测试修改效果,可以直接获取张量修改,但不会永久保存到数据集:
img, label = trainset[0] print(img[:,1,1]) img[:,1,1] = 0 print(img[:,1,1]) # 注意:这个修改后的img是临时变量,下次调用trainset[0]仍会返回原始数据
内容的提问来源于stack exchange,提问作者Winger 14
相关产品推荐
相关产品推荐

