大数跨境

Torchvision(中):数据增强,让数据更加多样性

Torchvision(中):数据增强,让数据更加多样性 知识代码AI
2026-08-24
2

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((200200), 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((6070))
random_crop_oper = transforms.RandomCrop((8080))
five_crop_oper = transforms.FiveCrop((6070))

# 原图
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 个:

  1. LinearTransformation(线性变换)
  2. Normalize(标准化)✅ 最常用
  3. RandomErasing(随机擦除)
  4. 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.50.50.5), (0.50.50.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((200200)),
    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.4850.4560.406], 
                         std=[0.2290.2240.225])
])

经验分享:数据增强的方法有很多,但并不是用得越多,效果越好。[pdf_7]


七、核心要点总结

操作类别
方法
关键参数
主要用途
类型转换
ToTensor()
PIL.Image/Numpy → Tensor
类型转换
ToPILImage(mode=None)
mode 根据维度推断
Tensor/Numpy → PIL.Image
尺寸调整
Resize(size, interpolation=2)
size 可为 tuple 或 int
调整图像尺寸
剪裁
CenterCrop(size)
size 可为 tuple 或 int
中心剪裁
剪裁
RandomCrop(size, padding=None)
padding 默认 None
随机位置剪裁
剪裁
FiveCrop(size)
size 可为 int 或 tuple
四角和中心共剪裁成五块
翻转
RandomHorizontalFlip(p=0.5)
p 为翻转概率
随机水平翻转
翻转
RandomVerticalFlip(p=0.5)
p 为翻转概率
随机垂直翻转
标准化
Normalize(mean, std, inplace=False)
各通道均值与标准差
数据标准化正则化
组合
Compose(transforms)
Transform 对象列表
多个变换连续操作

八、每课一练

问题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]

九、小结

  1. 数据类型转换ToTensor() 将 PIL/Numpy 转为 Tensor;ToPILImage() 将 Tensor/Numpy 转回 PIL.Image。
  2. 图像尺寸调整Resize 支持 tuple 精确匹配和 int 按比例缩放。
  3. 图像剪裁CenterCrop(中心)、RandomCrop(随机)、FiveCrop(四角+中心共五块)。
  4. 图像翻转RandomHorizontalFlip(水平翻转)、RandomVerticalFlip(垂直翻转),随机概率由 p 控制。
  5. 标准化Normalize 对 Tensor 进行标准化,加速收敛、提高训练效果。
  6. 变换组合Compose 将多个变换组合成 pipeline,按顺序执行。
  7. 与 datasets 结合:在 torchvision.datasets 的 transform 参数中传入 compose 后的操作,加载数据时自动完成预处理与增强。
  8. 核心原则:数据增强方法很多,但不是用得越多效果越好,需要根据实际任务合理选择。


【声明】内容源于网络
0
0
知识代码AI
技术基底 机器视觉全栈 × 光学成像 × 图像处理算法 编程栈 C++/C#工业开发 | Python智能建模 工具链 Halcon/VisionPro工业部署 | PyTorch/TensorFlow模型炼金术 | 模型压缩&嵌入式移植
内容 403
粉丝 0
知识代码AI 技术基底 机器视觉全栈 × 光学成像 × 图像处理算法 编程栈 C++/C#工业开发 | Python智能建模 工具链 Halcon/VisionPro工业部署 | PyTorch/TensorFlow模型炼金术 | 模型压缩&嵌入式移植
总阅读6.6k
粉丝0
内容403