如何使用PyTorch实现MNIST数据集的负片转换?
生成MNIST负片的PyTorch实现方案
没问题,我来帮你一步步解决生成MNIST负片的问题,先从你关心的数据类型说起,再讲两种常用的负片生成方法:
一、先搞清楚当前数据集的数据类型
你用了transforms.ToTensor(),这个操作会把MNIST原本0-255范围的uint8格式PIL图像,转换成0-1范围的torch.float32张量。你可以快速验证这一点:
# 取出一个样本查看 sample_image, sample_label = train_set[0] print(f"数据类型: {sample_image.dtype}") print(f"像素值范围: {sample_image.min()} ~ {sample_image.max()}")
运行后会输出torch.float32和0.0 ~ 1.0,这是PyTorch处理图像数据的标准格式,也方便后续模型计算。
二、生成负片的两种方法
因为像素已经被归一化到0-1,所以负片的计算逻辑是1 - 像素值(对应原始255范围的255 - 原始像素值),下面是两种可行的实现方式:
方法1:自定义Transform,加载数据集时直接生成负片
这种方法最优雅,相当于把负片处理整合到数据加载流程里,后续使用DataLoader时拿到的直接就是负片数据:
import torchvision as tv import torch from torchvision import transforms # 方式A:用Lambda Transform快速实现 transform = transforms.Compose([ transforms.ToTensor(), # 先转成0-1的float张量 transforms.Lambda(lambda x: 1 - x) # 生成负片 ]) # 方式B:自定义Transform类(更适合复杂逻辑扩展) class NegativeImageTransform: def __call__(self, tensor): return 1 - tensor # 替换成自定义类的写法 transform = transforms.Compose([ transforms.ToTensor(), NegativeImageTransform() ]) # 重新加载数据集 train_set = tv.datasets.MNIST(root="./data", train=True, download=True, transform=transform) test_set = tv.datasets.MNIST(root="./data", train=False, download=True, transform=transform) # 初始化DataLoader(和你原来的代码一致即可) train_dl = torch.utils.data.DataLoader(train_set, batch_size=64, shuffle=True) val_dl = torch.utils.data.DataLoader(test_set, batch_size=1000, shuffle=False)
方法2:迭代DataLoader时临时生成负片
如果你不想修改原始数据集,只是在使用时临时生成负片,可以在遍历DataLoader的时候处理:
# 保持你原来的数据集和DataLoader代码不变 for images, labels in train_dl: # 对当前batch的图像生成负片 negative_images = 1 - images # 接下来就可以用negative_images做训练、可视化等操作了
小提示:验证负片效果
如果想确认负片是否生成正确,可以用matplotlib可视化看看:
import matplotlib.pyplot as plt # 取一个原始样本和负片样本对比 original_image = tv.datasets.MNIST(root="./data", train=True, download=True, transform=transforms.ToTensor())[0][0] negative_image = 1 - original_image plt.subplot(1,2,1) plt.imshow(original_image.squeeze(), cmap="gray") plt.title("Original") plt.subplot(1,2,2) plt.imshow(negative_image.squeeze(), cmap="gray") plt.title("Negative") plt.show()
内容的提问来源于stack exchange,提问作者Michael
相关产品推荐
相关产品推荐

