MNIST 代码解析(3)构建pipline
2024-03-21 13:58:54
首先来看代码
#构造通道
pipeline = transforms.Compose([
transforms.ToTensor(),
transforms.Nomalize((0.1307,),(0.3081,))
])
这段代码使用了Pytorch库中的transforms模块定义了一个数据预处理流程,这个流程包含了两个步骤:ToTensor和Normalize。这种预处理流程通常用于图像数据的处理,特别是在训练神经网络模型的时候。
-
transforms.ToTensor()
功能:这个转换操作将PIL图像或NumPyndarray转换为torch.Tensor。在PyTorch中,Tensor是一个多维数组,用于存储和操作数据,类似于NumPy的ndarray,但它还可以在GPU上运行以加速计算。
细节:在转换过程中,图像的数据类型会从int8(取值范围0-255)转换为float32(浮点数),并且每个像素值会被标准化到[0,1]的范围内。这意味着原始的像素值255(表示白色)会变成1.0,像素值0(表示黑色)会变成0.0。 -
transforms.Normalize((0.1307,),(0.3081,))
功能:这个操作对张量图像进行标准化处理。标准化是将数据按比例缩放,使之落入一个小的特定区间。这在神经网络中非常重要,因为它有助于加快训练过程,减少模型初始化对模型训练结果的影响。
细节:
- 第一个参数
(0.1307,)是每个通道的均值。在这个例子中,我们假设图像是灰度的(只有一个通道),所以只有一个元素。这个值是用于将图像数据的每个通道中心化(即减去均值)。 - 第二个参数
(0,3081,)是每个通道的标准差。这同样假设图像是灰度的。这个值用于缩放图像数据,使得每个通道的数据分布具有单位标准差。 - 通过这种方式,图像数据会被进一步标准化,有助于模型学习和泛化。
如果图像不是灰度的,会有什么变化?
如果图像不是灰度的,而是彩色的(通常是RGB三通道)
transforms.Normalize操作的参数需要相应地调整以适应每个通道。对于RGB图像,你需要为每个通道指定一个标准差和均值,于是现在参数将是三个元素的元组,而不是单一元素的元组。
彩色图像的transforms.Normalize
对于彩色图像,transforms.Normalize的调用可能看起来像这样:
transforms.Noramlize((mean_R, mean_G, mean_B),(std_R, std_G, std_B))
可以知道的是:
(mean_R, mean_G, mean_B)是RGB三个通道的均值。这些值用于将每个通道的数据中心化。
(std_R, std_G, std_B)是RGB三个通道的标准差。这些值用于缩放每个通道的数据,使得每个通道的数据分布具有单位标准差。

浙公网安备 33010602011771号