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

如何使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.29 13:32:35