Torchvision(中):数据增强,让数据更加多样性
一、模块概述
torchvision.transforms 是 Torchvision 库中提供常用图像操作的工具包,支持对 Tensor 及 PIL Image 对象的操作,功能包括:随机切割、旋转、数据类型转换等。
按照功能可分为以下几类:
-
数据类型转换 -
对 PIL.Image 和 Tensor 进行变换 -
变换的组合[pdf_7]
二、数据类型转换
1. transforms.ToTensor()
-
作用:将 PIL.Image 或 Numpy.ndarray 格式的数据转化为 Tensor 格式 -
使用场景:模型训练阶段需要传入 Tensor 类型的数据,而读取的数据集图片是 PIL.Image 对象
2. transforms.ToPILImage(mode=None)
-
作用:将 Tensor 或 Numpy.ndarray 格式的数据转化为 PIL.Image 格式(ToTensor 的逆操作) -
mode 参数:代表 PIL.Image 的模式,默认值为 None,根据输入数据的维度进行推断: -
输入为 3 通道 → mode 为 'RGB' -
输入为 4 通道 → mode 为 'RGBA' -
输入为 2 通道 → mode 为 'LA' -
输入为单通道 → mode 根据输入数据的类型确定具体模式
代码示例
from PIL import Image
from torchvision import transforms
img = Image.open('jk.jpg')
print(type(img)) # <class 'PIL.JpegImagePlugin.JpegImageFile'>
# PIL.Image 转换为 Tensor
img1 = transforms.ToTensor()(img)
print(type(img1)) # <class 'torch.Tensor'>
# Tensor 转换为 PIL.Image
img2 = transforms.ToPILImage()(img1)
print(type(img2)) # <class 'PIL.Image.Image'>
★说明:PIL.JpegImagePlugin.JpegImageFile 类是 PIL.Image.Image 类的子类。[pdf_7]
三、对 PIL.Image 和 Tensor 进行变换
这些变换操作可接收多种数据格式,可直接对 PIL 格式或 Tensor 进行变换,无需额外做数据类型转换。
1. Resize(尺寸调整)
定义:
torchvision.transforms.Resize(size, interpolation=2)
参数说明:
-
size:期望输出的尺寸 -
如果是元组 (h, w):图像输出尺寸将与之匹配 -
如果是 int 类型整数:图像较小的边将被匹配到该整数,另一条边按比例缩放 -
interpolation:插值算法,int 类型,默认为 2,表示 PIL.Image.BILINEAR
代码示例:
from PIL import Image
from torchvision import transforms
# 定义一个 Resize 操作
resize_img_oper = transforms.Resize((200, 200), interpolation=2)
# 原图
orig_img = Image.open('jk.jpg')
# Resize 操作后的图
img = resize_img_oper(orig_img)
★注意:训练时通常要把图片 resize 到一定大小(如 128×128、256×256)。如果设定为 int 型,较长的边会按比例缩放。resize 之后一般会接 crop 操作,但高与宽差距较大时,会 crop 掉很多有用信息。[pdf_7]
2. 剪裁(Cropping)
torchvision.transforms 提供多种剪裁方法:中心剪裁、随机剪裁、四角和中心剪裁等。
(1)CenterCrop(中心剪裁)
torchvision.transforms.CenterCrop(size)
-
size:期望输出的剪裁尺寸 -
元组 (h, w):剪裁后的图像尺寸将与之匹配 -
int 类型:剪裁出来的图像是 (size, size)的正方形
(2)RandomCrop(随机剪裁)
torchvision.transforms.RandomCrop(size, padding=None)
-
size:期望输出的剪裁尺寸,用法同上 -
padding:图像每个边框上的可选填充,默认值为 None(即没有填充),通常不会使用
(3)FiveCrop(四角和中心剪裁)
torchvision.transforms.FiveCrop(size)
-
将给定的 PIL Image 或 Tensor,分别从四角和中心进行剪裁,共剪裁成五块 -
size 可以是 int 或 tuple,用法同上
代码示例:
from PIL import Image
from torchvision import transforms
# 定义剪裁操作
center_crop_oper = transforms.CenterCrop((60, 70))
random_crop_oper = transforms.RandomCrop((80, 80))
five_crop_oper = transforms.FiveCrop((60, 70))
# 原图
orig_img = Image.open('jk.jpg')
# 中心剪裁
img1 = center_crop_oper(orig_img)
# 随机剪裁
img2 = random_crop_oper(orig_img)
# 四角和中心剪裁
imgs = five_crop_oper(orig_img)
for img in imgs:
display(img)
3. 翻转(Flipping)
(1)RandomHorizontalFlip(随机水平翻转)
torchvision.transforms.RandomHorizontalFlip(p=0.5)
-
p:随机翻转的概率值,默认为 0.5 -
如果想要必须执行翻转操作,将 p 设置为 1 即可
(2)RandomVerticalFlip(随机垂直翻转)
torchvision.transforms.RandomVerticalFlip(p=0.5)
-
p:随机翻转的概率值,默认为 0.5
代码示例:
from PIL import Image
from torchvision import transforms
# 定义翻转操作
h_flip_oper = transforms.RandomHorizontalFlip(p=1)
v_flip_oper = transforms.RandomVerticalFlip(p=1)
# 原图
orig_img = Image.open('jk.jpg')
# 水平翻转
img1 = h_flip_oper(orig_img)
# 垂直翻转
img2 = v_flip_oper(orig_img)
四、只对 Tensor 进行变换
目前版本的 Torchvision(v0.10.0)对各种图像变换操作已基本同时支持 PIL Image 和 Tensor 类型,只针对 Tensor 的变换操作仅有 4 个:
-
LinearTransformation(线性变换) -
Normalize(标准化)✅ 最常用 -
RandomErasing(随机擦除) -
ConvertImageDtype(格式转换)
标准化 Normalize
数学公式:
output = (input - mean) / std
目的:对图像的每个通道利用均值和标准差进行正则化,保证数据集中所有图像分布相似,训练时更容易收敛,既加快训练速度,也提高训练效果。
定义:
torchvision.transforms.Normalize(mean, std, inplace=False)
参数说明:
-
mean:各通道的均值 -
std:各通道的标准差 -
inplace:是否原地操作,默认为 False
代码示例:
from PIL import Image
from torchvision import transforms
# 定义标准化操作
norm_oper = transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
# 原图
orig_img = Image.open('jk.jpg')
# 图像转化为 Tensor
img_tensor = transforms.ToTensor()(orig_img)
# 标准化
tensor_norm = norm_oper(img_tensor)
# Tensor 转化为图像
img_norm = transforms.ToPILImage()(tensor_norm)
★说明:标准化是一个常规做法,无脑进行标准化后再训练的效果,大概率要好于不进行标准化。标准化会将数据映射到同一区间中,同一类别的图片像素值可能有差异,但其分布都是类似的分布。[pdf_7]
五、变换的组合(Compose)
作用:将多个变换组合到一起,进行连续操作。
定义:
torchvision.transforms.Compose(transforms)
参数说明:
-
transforms:一个 Transform 对象的列表,表示要组合的变换列表
代码示例(将图片变为 200×200 像素大小,并随机裁切成 80 像素正方形):
from PIL import Image
from torchvision import transforms
# 原图
orig_img = Image.open('jk.jpg')
# 定义组合操作
composed = transforms.Compose([
transforms.Resize((200, 200)),
transforms.RandomCrop(80)
])
# 组合操作后的图
img = composed(orig_img)
六、结合 datasets 使用
在利用 torchvision.datasets 读取数据集时,可通过 transform 参数对图像进行预处理操作(数据增强、归一化、旋转或缩放等)。该参数可接收一个 torchvision.transforms 操作或由 Compose 类定义的操作组合。
代码示例(读取 MNIST 数据集):
from torchvision import transforms
from torchvision import datasets
# 定义一个 transform
my_transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5), (0.5))
])
# 读取 MNIST 数据集,同时做数据变换
mnist_dataset = datasets.MNIST(
root='./data',
train=False,
transform=my_transform,
target_transform=None,
download=True
)
# 查看变换后的数据类型
item = mnist_dataset.__getitem__(0)
print(type(item[0]))
# 输出:<class 'torch.Tensor'>
实际项目中的 transform 示例(图像分类实战):
transform = transforms.Compose([
transforms.RandomResizedCrop(dest_image_size),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225])
])
★经验分享:数据增强的方法有很多,但并不是用得越多,效果越好。[pdf_7]
七、核心要点总结
|
|
|
|
|
|---|---|---|---|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
八、每课一练
问题:transforms.ToTensor() 和 transforms.Normalize() 分别有什么作用?在 pipeline 中它们的顺序是怎样的?
解答:
-
transforms.ToTensor():将 PIL.Image 或 Numpy.ndarray 格式的数据转化为 Tensor 格式 -
transforms.Normalize():对图像的每个通道利用均值和标准差进行标准化,公式为output = (input - mean) / std -
顺序:先 ToTensor()将数据转为 Tensor,再Normalize()进行标准化,因为 Normalize 只对 Tensor 类型的数据进行操作[pdf_7]
九、小结
-
数据类型转换: ToTensor()将 PIL/Numpy 转为 Tensor;ToPILImage()将 Tensor/Numpy 转回 PIL.Image。 -
图像尺寸调整: Resize支持 tuple 精确匹配和 int 按比例缩放。 -
图像剪裁: CenterCrop(中心)、RandomCrop(随机)、FiveCrop(四角+中心共五块)。 -
图像翻转: RandomHorizontalFlip(水平翻转)、RandomVerticalFlip(垂直翻转),随机概率由 p 控制。 -
标准化: Normalize对 Tensor 进行标准化,加速收敛、提高训练效果。 -
变换组合: Compose将多个变换组合成 pipeline,按顺序执行。 -
与 datasets 结合:在 torchvision.datasets的 transform 参数中传入 compose 后的操作,加载数据时自动完成预处理与增强。 -
核心原则:数据增强方法很多,但不是用得越多效果越好,需要根据实际任务合理选择。

